Lesson 51: Custom Loss Functions & Autograd

USAAIO Lesson 51, from Week 18 of Phase 3, on implementing custom differentiable losses in PyTorch. It covers focal loss for class imbalance, triplet loss with online hard-negative mining for metric learning, contrastive loss as the foundation of CLIP, and torch.autograd.Function for gradients that are not standard. Every formula was verified by real execution with torch 2.7.1 and sklearn in June 2026. The lesson runs to 30 slides.

Subject: Machine Learning · 59 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. Custom Loss Functions & Autograd

Title

USAAIO · Lesson 51 · Week 18 (Phase 3)

When cross-entropy isn't enough: focal loss for imbalance, triplet and contrastive losses for metric learning, and writing your own differentiable op with torch.autograd.Function.

2. By the end of this lesson you can

Objectives

  1. Implement focal loss from scratch and verify its gradient
  2. Explain why focal loss down-weights easy examples and when to use it
  3. Implement triplet loss with online hard-negative mining
  4. Distinguish contrastive loss from triplet loss and state how it underlies CLIP
  5. Write a custom torch.autograd.Function with correct forward and backward methods

3. What survived from Gradient Boosting?

Warm-up

Discussion prompt

Before we open Lesson 51: Custom Loss Functions & Autograd: without looking back, what was the main idea of Gradient Boosting, and what could you do by the end of it that you could not do before?

Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.

Answer:

the gradient boosting framework F_m = F_{m-1} + eta*h_m, pseudo-residuals as negative gradients of the loss (y-F for MSE; y-sigma(F) for classification), a from-scratch 3-round GBM regressor using sklearn DecisionTree, XGBoost improvements (L2 leaf regularization, column subsampling, approximate splits), LightGBM histogram-based leaf-wise growth, and a cross-validated hyperparameter sweep. Build GradientBoostingRegressorFromScratch on make_regression data and compare GBM vs Random Forest.

4. The differentiability requirement

Section

Part 1 of 4

5. What makes a custom loss valid?

Concept

Any Python function can compute a number. But only functions built from torch operations are automatically differentiable — PyTorch's autograd engine tracks the computation graph.

Verify with torch.autograd.gradcheck: numerically differentiates your function and compares it to autograd. A passing gradcheck means your gradients are correct.

6. Break it if you can: What makes a custom loss valid?

Counterexample

Discussion prompt

Any Python function can compute a number. But only functions built from torch operations are automatically differentiable — PyTorch's autograd engine tracks the computation graph.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Verify with torch.autograd.gradcheck: numerically differentiates your function and compares it to autograd. A passing gradcheck means your gradients are correct.

7. Focal loss: down-weight the easy ones

Section

Part 2 of 4

8. The imbalance problem

Concept

Object detectors evaluate ~100 000 anchor boxes per image. Fewer than 10 contain objects. Cross-entropy loss is dominated by the overwhelming majority of easy negatives — the model learns nothing about the rare positives.

classcountrole in BCE loss
background (easy neg)~99 900dominates gradient
object (positive)~100signal drowned out

RetinaNet's fix (Lin et al., ICCV 2017): add a modulating factor that shrinks the loss on easy examples exponentially with a focusing parameter gamma.

9. Fill in: count for The imbalance problem

Comparison

Comparison matrix

From The imbalance problem: refill the count column from what you know. The rest of the table is as it appeared.

classcountrole in BCE loss
background (easy neg)~99 900dominates gradient
object (positive)~100signal drowned out

10. Focal loss formula

Concept

\[ \text{FL}(p_t) = -(1-p_t)^\gamma \log(p_t) \]

where p_t = p when the target is 1 (positive), p_t = 1-p when the target is 0 (negative). p = sigmoid(logit).

p (confidence)gamma=0 (CE)gamma=1gamma=2
0.1 (hard)2.30262.07231.8651
0.5 (medium)0.69310.34660.1733
0.9 (easy)0.10540.01050.0011

At gamma=2 the easy example (p=0.90) is downweighted by factor 0.002 versus gamma=0. Hard examples (p=0.10) lose only 19% — they dominate training.

11. What each one costs: Focal loss formula

Trade off

Comparison matrix

From Focal loss formula: every row here is a choice with a cost. Fill the gamma=2 column, then say which row you would actually pick and what you give up for it.

p (confidence)gamma=0 (CE)gamma=1gamma=2
0.1 (hard)2.30262.07231.8651
0.5 (medium)0.69310.34660.1733
0.9 (easy)0.10540.01050.0011

12. Guess the shape of the answer: Focal loss from scratch

Estimation

Predict first

Implement focal_loss using only torch.sigmoid and F.binary_cross_entropy_with_logits. Verify on a 4-sample batch.

Commit before you compute: what does Focal loss from scratch come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: mean focal loss = 0.2437; gradient flows back through the computation graph

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Built entirely from torch ops — autograd tracks the graph automatically.

13. Focal loss from scratch

Worked example

Implement focal_loss using only torch.sigmoid and F.binary_cross_entropy_with_logits. Verify on a 4-sample batch.

import torch, torch.nn.functional as F

def focal_loss(logits, targets, gamma=2.0):
    p    = torch.sigmoid(logits)
    bce  = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')
    p_t  = p * targets + (1 - p) * (1 - targets)  # p_t matches label
    return ((1 - p_t) ** gamma * bce).mean()

logits  = torch.tensor([3.0, 0.1, -1.0, 0.5], requires_grad=True)
targets = torch.tensor([1.0, 0.0,  1.0, 1.0])
loss = focal_loss(logits, targets, gamma=2)
print(round(loss.item(), 4))   # 0.2437
loss.backward()
print(logits.grad.round(decimals=4))

mean focal loss = 0.2437; gradient flows back through the computation graph

Why: Built entirely from torch ops — autograd tracks the graph automatically. gradcheck also passes (verified separately).

logittargetp(1-p_t)^2BCEFL
3.010.95260.00220.04860.0001
0.100.52500.27560.74440.2052
-1.010.26890.53441.31330.7019
0.510.62250.14250.47410.0676

14. Which is which, by target

Discrimination

Sort into buckets

Sort these by target, from memory, without looking back at Focal loss from scratch. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

1
3.0; -1.0; 0.5
0
0.1
g1
target is "1" for 3.0, -1.0, 0.5 — that is what the table on "Focal loss from scratch" records, and it is the single property separating this group from the rest.
g2
target is "0" for 0.1 — that is what the table on "Focal loss from scratch" records, and it is the single property separating this group from the rest.

15. CE vs focal on an imbalanced dataset

Concept

10% positive class (50 positives, 450 negatives, make_classification). A 2-layer MLP trained for 100 epochs with Adam (lr=0.01).

lossoverall accminority recall
cross-entropy0.9480.560
focal (gamma=2)0.9600.680

The accuracy improvement is modest (0.948 → 0.960), but minority recall jumps from 56% to 68% — the class that matters most in imbalanced tasks.

16. By analogy: CE vs focal on an imbalanced dataset

Analogy

Discussion prompt

Explain CE vs focal on an imbalanced dataset by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

The accuracy improvement is modest (0.948 → 0.960), but minority recall jumps from 56% to 68% — the class that matters most in imbalanced tasks.

17. Something is wrong here: using Python math inside a custom loss

Anomaly

Predict first

A student writes this, and it looks reasonable:

Compute p_t with math.log since it's just a scalar operation.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: .item() detaches the tensor from the autograd graph.

Keep everything in torch operations — never call .item() inside the loss.

Why: .item() detaches the tensor from the autograd graph. math.log returns a plain Python float — no gradient is tracked.

18. Trap: using Python math inside a custom loss

Trap

The trap

Compute p_t with math.log since it's just a scalar operation.

import math; loss = -(1-p_t.item())**gamma * math.log(p_t.item())

Why: .item() detaches the tensor from the autograd graph. math.log returns a plain Python float — no gradient is tracked.

loss.backward() raises AttributeError: 'float' has no attribute 'backward'

Why: The loss is no longer a tensor at all, let alone a leaf in the computation graph.

The fix

Keep everything in torch operations — never call .item() inside the loss.

loss = ((1 - p_t) ** gamma * bce).mean()

Why: All ops are torch.*. The computation graph is intact. .backward() works and fills .grad on every tensor that has requires_grad=True.

Verify once: torch.autograd.gradcheck(focal_loss, (x_f64, t)) should print True

Why: gradcheck numerically estimates Jacobians and compares them to autograd — the gold standard for custom loss correctness.

19. Break it on purpose: using Python math inside a custom loss

Break the constraint

Discussion prompt

The rule this trap just fixed:

All ops are torch.*. The computation graph is intact. .backward() works and fills .grad on every tensor that has requires_grad=True.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

.item() detaches the tensor from the autograd graph. math.log returns a plain Python float — no gradient is tracked.

20. Triplet loss: metric learning

Section

Part 3 of 4

21. The triplet loss idea

Concept

Instead of predicting a label, the network learns an embedding space where same-class points are close and different-class points are far — useful for face recognition and image retrieval.

\[ \mathcal{L} = \max\!\bigl(0,\; d(a,p) - d(a,n) + \text{margin}\bigr) \]

The anchor a, positive p (same class), and negative n (different class) form a triplet. The loss is zero when the negative is already margin further away than the positive.

22. Teach it back: The triplet loss idea

Explain it

Discussion prompt

Explain The triplet loss idea to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

Instead of predicting a label, the network learns an embedding space where same-class points are close and different-class points are far — useful for face recognition and image retrieval.

23. Guess the shape of the answer: Triplet loss: easy vs hard negative

Estimation

Predict first

2-D embeddings. Anchor at origin. Compare an easy (far) negative vs a hard (close) negative with the same positive.

Commit before you compute: what does Triplet loss: easy vs hard negative come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Easy negative: loss = max(0, 0.5831 - 2.5000 + 1.0) = max(0, -0.9169) = 0.0000

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Already well-separated — no gradient, no update.

24. Triplet loss: easy vs hard negative

Worked example

2-D embeddings. Anchor at origin. Compare an easy (far) negative vs a hard (close) negative with the same positive.

import torch

def triplet_loss(anchor, pos, neg, margin=1.0):
    d_pos = torch.norm(anchor - pos,  dim=1)
    d_neg = torch.norm(anchor - neg, dim=1)
    return torch.clamp(d_pos - d_neg + margin, min=0.0).mean()

a = torch.tensor([[0.0, 0.0]])
p = torch.tensor([[0.5, 0.3]])   # d(a,p)=0.5831
n_easy = torch.tensor([[2.0, 1.5]])  # d(a,n)=2.5000
n_hard = torch.tensor([[0.6, 0.4]])  # d(a,n)=0.7211

print(triplet_loss(a, p, n_easy).item())  # 0.0000
print(triplet_loss(a, p, n_hard).item())  # 0.8620

Easy negative: loss = max(0, 0.5831 - 2.5000 + 1.0) = max(0, -0.9169) = 0.0000

Why: Already well-separated — no gradient, no update. This is the 'easy triplet' problem that motivates hard negative mining.

scenariod(a,p)d(a,n)loss
easy negative (far)0.58312.50000.0000
hard negative (close)0.58310.72110.8620

25. Work backwards from the answer: Triplet loss: easy vs hard negative

Reverse engineer

Discussion prompt

Work backwards. The example finished here:

Easy negative: loss = max(0, 0.5831 - 2.5000 + 1.0) = max(0, -0.9169) = 0.0000

What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.

Hint: Every quantity in the result had to enter somewhere. Account for each one.

Answer:

2-D embeddings. Anchor at origin. Compare an easy (far) negative vs a hard (close) negative with the same positive.

26. Online hard negative mining

Concept

Pre-selecting hard triplets offline is expensive and goes stale. Online hard mining picks the hardest triplet within each mini-batch on every step.

  1. Compute all pairwise distances in the batch
  2. For each anchor: find the hardest positive (same label, max distance)
  3. For each anchor: find the hardest negative (diff label, min distance)
  4. Form triplets from these extremes; average the non-zero losses

On a batch of 4 embeddings (dim=8, labels [0,0,1,1]) the hard triplet loss is 1.8210 — nearly every pair violates the margin, providing constant gradient.

27. Teach it back: Online hard negative mining

Explain it

Discussion prompt

Explain Online hard negative mining to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

Pre-selecting hard triplets offline is expensive and goes stale. Online hard mining picks the hardest triplet within each mini-batch on every step.

28. Contrastive loss & CLIP

Concept

\[ \mathcal{L}_{\text{contrastive}} = y\,d^2 + (1-y)\,\max(0,\, m - d)^2 \]

Label y=1 means same class (pull together, minimize d²). Label y=0 means different class (push apart if d < margin). Unlike triplet loss, it operates on pairs not triplets.

pair typedistance dloss (m=1.0)
same class, d=0.20.200.0400 (= d²)
same class, d=0.80.800.6400 (= d²)
diff class, d=0.20.200.6400 (= (1-0.2)²)
diff class, d=1.51.500.0000 (already separated)

CLIP (Radford et al. 2021) extends this to image-text pairs at scale, learning a joint embedding space where a dog photo and the text 'a dog' are near each other.

29. Fill in: loss (m=1.0) for Contrastive loss & CLIP

Comparison

Comparison matrix

From Contrastive loss & CLIP: refill the loss (m=1.0) column from what you know. The rest of the table is as it appeared.

pair typedistance dloss (m=1.0)
same class, d=0.20.200.0400 (= d²)
same class, d=0.80.800.6400 (= d²)
diff class, d=0.20.200.6400 (= (1-0.2)²)
diff class, d=1.51.500.0000 (already separated)

30. Something is wrong here: wrong p_t in focal loss

Anomaly

Predict first

A student writes this, and it looks reasonable:

Compute p_t = p for all samples (regardless of target label).

It is wrong. Say what breaks — and say it before you turn the page.

Correct: For target=0 (negative), a high-confidence prediction gives p near 1 — but the modulating factor (1-p)^2 would be near 0, wrongly suppressing the negative's gradient.

Match p_t to the label: it is the probability assigned to the correct class.

Why: For target=0 (negative), a high-confidence prediction gives p near 1 — but the modulating factor (1-p)^2 would be near 0, wrongly suppressing the negative's gradient.

31. Trap: wrong p_t in focal loss

Trap

The trap

Compute p_t = p for all samples (regardless of target label).

p_t = torch.sigmoid(logits) # always p

Why: For target=0 (negative), a high-confidence prediction gives p near 1 — but the modulating factor (1-p)^2 would be near 0, wrongly suppressing the negative's gradient.

The model learns only from positive examples; easy negatives are never penalized correctly

Why: The original focal loss paper explicitly defines p_t = p when y=1 and p_t = 1-p when y=0, so the factor always reflects how confident-and-correct the prediction is.

The fix

Match p_t to the label: it is the probability assigned to the correct class.

p_t = p * targets + (1 - p) * (1 - targets)

Why: When target=1, p_t=p; when target=0, p_t=1-p. A high-confidence correct prediction (any label) has p_t near 1 and small (1-p_t)^2 — correctly down-weighted.

Verify: for logit=3, target=0 (easy neg), p=0.95, p_t=0.05, (1-p_t)^2=0.9025 — NOT suppressed

Why: The model is confidently wrong here — it should be heavily penalized, and p_t=0.05 gives a large modulating factor.

32. Which of these survive contact with Lesson 51: Custom Loss Functions & Autograd?

Two truths and a lie

Sort into buckets

Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.

Holds up
Verify with torch.autograd.gradcheck: numerically differentiates your function and compares it to autograd. A passing gradcheck means your gradients are correct.; RetinaNet's fix (Lin et al., ICCV 2017): add a modulating factor that shrinks the loss on easy examples exponentially with a focusing parameter gamma.; At gamma=2 the easy example (p=0.90) is downweighted by factor 0.002 versus gamma=0. Hard examples (p=0.10) lose only 19% — they dominate training.
Breaks
Compute p_t with math.log since it's just a scalar operation.; Compute p_t = p for all samples (regardless of target label).
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 51: Custom Loss Functions & Autograd puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

33. Custom `autograd.Function`

Section

Part 4 of 4

34. When built-in autograd isn't enough

Concept

Sometimes you have an operation whose forward pass can be computed efficiently, but autograd's automatic differentiation would be slow, numerically unstable, or simply wrong.

torch.autograd.Function lets you provide both the forward computation and the exact backward (gradient) formula yourself.

35. By analogy: When built-in autograd isn't enough

Analogy

Discussion prompt

Explain When built-in autograd isn't enough by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

Sometimes you have an operation whose forward pass can be computed efficiently, but autograd's automatic differentiation would be slow, numerically unstable, or simply wrong.

36. Guess the shape of the answer: Custom autograd: Swish activation

Estimation

Predict first

Implement Swish, f(x) = x * sigmoid(x), as a custom autograd Function. The closed-form gradient is sigma(x) + x * sigma(x) * (1 - sigma(x)).

Commit before you compute: what does Custom autograd: Swish activation come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: At x=1.0: forward = 0.7311; gradient = 0.7311 + 1.00.73110.2689 = 0.7311 + 0.1966 = 0.9277

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. ctx.save_for_backward caches tensors needed in backward without keeping the full forward graph.

37. Custom autograd: Swish activation

Worked example

Implement Swish, f(x) = x * sigmoid(x), as a custom autograd Function. The closed-form gradient is sigma(x) + x * sigma(x) * (1 - sigma(x)).

import torch

class Swish(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x):
        sig = torch.sigmoid(x)
        ctx.save_for_backward(x, sig)   # cache for backward
        return x * sig

    @staticmethod
    def backward(ctx, grad_output):
        x, sig = ctx.saved_tensors
        grad = sig + x * sig * (1 - sig)   # d/dx [x*sigma(x)]
        return grad_output * grad

x = torch.tensor([-1.0, 0.0, 1.0, 2.0], requires_grad=True)
out = Swish.apply(x)
out.sum().backward()
print(x.grad.round(decimals=4))

At x=1.0: forward = 0.7311; gradient = 0.7311 + 1.00.73110.2689 = 0.7311 + 0.1966 = 0.9277

Why: ctx.save_for_backward caches tensors needed in backward without keeping the full forward graph. The chain rule is applied via grad_output * df/dx.

xsigma(x)forward f(x)grad df/dx
-1.00.2689-0.26890.0723
0.00.50000.00000.5000
1.00.73110.73110.9277
2.00.88081.76161.0918

38. Watch it run: Custom autograd: Swish activation

Pattern

Step through it

Step through Custom autograd: Swish activation one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: x is -1.0
  2. Step 2: x is 0.0
  3. Step 3: x is 1.0
  4. Step 4: x is 2.0

39. Rebuild the recipe: Custom loss & autograd recipe

Ranking

Put in order

These are the steps of Custom loss & autograd recipe, scrambled. Put them back in order before the next slide shows you.

  1. Focal loss: p_t = p*y + (1-p)*(1-y) → FL = ((1-p_t)**gamma * BCE).mean()
  2. Triplet loss: clamp(d(a,p) - d(a,n) + margin, min=0); mine hard negatives online
  3. Contrastive loss: y*d^2 + (1-y)*clamp(m-d, min=0)^2; operates on pairs
  4. Custom Function: forward returns result + calls ctx.save_for_backward; backward returns grad_output * df/dx
  5. Verify always: torch.autograd.gradcheck with float64 inputs before shipping

Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.

40. Custom loss & autograd recipe

Pattern

  1. Focal loss: p_t = p*y + (1-p)*(1-y) → FL = ((1-p_t)**gamma * BCE).mean()
  2. Triplet loss: clamp(d(a,p) - d(a,n) + margin, min=0); mine hard negatives online
  3. Contrastive loss: y*d^2 + (1-y)*clamp(m-d, min=0)^2; operates on pairs
  4. Custom Function: forward returns result + calls ctx.save_for_backward; backward returns grad_output * df/dx
  5. Verify always: torch.autograd.gradcheck with float64 inputs before shipping

41. Where does each piece belong: Lesson 51: Custom Loss Functions & Autograd

Sorting

Sort into buckets

These are the pieces of Lesson 51: Custom Loss Functions & Autograd, out of order. Put each one back under the part of the lesson it belongs to.

Focal loss: down-weight the easy ones
The imbalance problem; Focal loss formula; Focal loss from scratch
Triplet loss: metric learning
The triplet loss idea; Triplet loss: easy vs hard negative; Online hard negative mining
Custom autograd.Function
When built-in autograd isn't enough; Custom autograd: Swish activation; Custom loss & autograd recipe
s1
Focal loss: down-weight the easy ones is where Lesson 51: Custom Loss Functions & Autograd puts The imbalance problem, Focal loss formula, Focal loss from scratch. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
Triplet loss: metric learning is where Lesson 51: Custom Loss Functions & Autograd puts The triplet loss idea, Triplet loss: easy vs hard negative, Online hard negative mining. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Custom autograd.Function is where Lesson 51: Custom Loss Functions & Autograd puts When built-in autograd isn't enough, Custom autograd: Swish activation, Custom loss & autograd recipe. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

42. Rule out three: Check yourself — focal loss

Elimination

Eliminate the wrong options

A sample has logit=3.0, target=1, giving p=0.9526. With gamma=2, the focal loss is approximately:

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. 0.0001 — modulated down by (1-0.9526)^2 = 0.0022
  • B. 0.0486 — same as binary cross-entropy
  • C. 0.9526 — the probability itself
  • D. 0.2237 — the square root of p

Survives elimination: A

Why: BCE = 0.0486. Modulating factor = (1-0.9526)^2 = 0.0022. FL = 0.0022 * 0.0486 = 0.0001. This easy example contributes almost nothing to the gradient.

43. Check yourself — focal loss

Check

Compute before choosing.

Check your understanding

A sample has logit=3.0, target=1, giving p=0.9526. With gamma=2, the focal loss is approximately:

  • A. 0.0001 — modulated down by (1-0.9526)^2 = 0.0022 (correct)
  • B. 0.0486 — same as binary cross-entropy
  • C. 0.9526 — the probability itself
  • D. 0.2237 — the square root of p

Answer: A

Why: BCE = 0.0486. Modulating factor = (1-0.9526)^2 = 0.0022. FL = 0.0022 * 0.0486 = 0.0001. This easy example contributes almost nothing to the gradient.

Why B tempts people
That is the plain BCE. Focal loss multiplies BCE by (1-p_t)^gamma, shrinking it dramatically for well-classified examples.
Why C tempts people
p=0.9526 is the sigmoid output, not a loss. Losses decrease as confidence increases.
Why D tempts people
sqrt(p) has no role in focal loss. The modulating factor is (1-p_t)^gamma, not a root.

44. Answer it before you see the options: Check yourself — triplet loss

Prediction

Predict first

Anchor=a, positive=p, negative=n. d(a,p)=0.8, d(a,n)=1.2, margin=1.0. What is the triplet loss?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: 0.6 — from max(0, 0.8 - 1.2 + 1.0)

Why: Triplet loss = max(0, d(a,p) - d(a,n) + margin) = max(0, 0.8 - 1.2 + 1.0) = max(0, 0.6) = 0.6. The negative is farther than the positive, but not by a full margin — the constraint is violated.

45. Check yourself — triplet loss

Check

Work through the margin.

Check your understanding

Anchor=a, positive=p, negative=n. d(a,p)=0.8, d(a,n)=1.2, margin=1.0. What is the triplet loss?

  • A. 0.6 — from max(0, 0.8 - 1.2 + 1.0) (correct)
  • B. 0.0 — the negative is already farther away
  • C. 1.0 — equal to the margin
  • D. 0.4 — the difference d(a,n) - d(a,p)

Answer: A

Why: Triplet loss = max(0, d(a,p) - d(a,n) + margin) = max(0, 0.8 - 1.2 + 1.0) = max(0, 0.6) = 0.6. The negative is farther than the positive, but not by a full margin — the constraint is violated.

Why B tempts people
Loss=0 requires d(a,n) >= d(a,p) + margin = 0.8 + 1.0 = 1.8. Here d(a,n)=1.2, so the triplet constraint is still violated.
Why C tempts people
The margin is the required gap, not the loss value. The loss equals the shortfall: 0.6 short of the required 1.0 gap.
Why D tempts people
d(a,n) - d(a,p) = 0.4 measures the existing gap, but triplet loss subtracts this from the margin: margin - gap = 1.0 - 0.4 = 0.6.

46. Rule out three: Check yourself — custom autograd

Elimination

Eliminate the wrong options

Inside a torch.autograd.Function's backward method, how do you retrieve tensors saved during forward?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. ctx.saved_tensors — unpacks whatever was saved with ctx.save_for_backward
  • B. ctx.inputs — PyTorch stores all inputs automatically
  • C. self.cache — assign in forward, retrieve in backward via self
  • D. grad_output.saved_tensors — the gradient carries the cache

Survives elimination: A

Why: ctx.save_for_backward(x, sig) in forward; x, sig = ctx.saved_tensors in backward. This API ensures memory is freed correctly and works with gradient checkpointing.

47. Check yourself — custom autograd

Check

Which call does the right thing?

Check your understanding

Inside a torch.autograd.Function's backward method, how do you retrieve tensors saved during forward?

  • A. ctx.saved_tensors — unpacks whatever was saved with ctx.save_for_backward (correct)
  • B. ctx.inputs — PyTorch stores all inputs automatically
  • C. self.cache — assign in forward, retrieve in backward via self
  • D. grad_output.saved_tensors — the gradient carries the cache

Answer: A

Why: ctx.save_for_backward(x, sig) in forward; x, sig = ctx.saved_tensors in backward. This API ensures memory is freed correctly and works with gradient checkpointing.

Why B tempts people
ctx.inputs does not exist. PyTorch does not automatically store all inputs — you must explicitly save what you need.
Why C tempts people
torch.autograd.Function methods are static; there is no self. Using instance attributes would break thread-safety and gradient checkpointing.
Why D tempts people
grad_output is the upstream gradient tensor; it has no saved_tensors attribute. Tensors are saved on ctx, not on the gradient.

48. Your turn: build the losses

Section

Project

49. Project: custom losses end-to-end

Concept

Implement all three losses from scratch, verify gradients, and compare CE vs focal on an imbalanced dataset.

#milestonetool
1focal loss + gradchecksigmoid, BCE, autograd
2triplet loss + hard miningcdist, clamp, label masks
3CE vs focal on imbalanced datamake_classification, Adam

Build rules: implement each loss as a standalone function; pass float64 tensors to gradcheck; never call .item() inside the loss.

50. Break it if you can: Project: custom losses end-to-end

Counterexample

Discussion prompt

Implement all three losses from scratch, verify gradients, and compare CE vs focal on an imbalanced dataset.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Build rules: implement each loss as a standalone function; pass float64 tensors to gradcheck; never call .item() inside the loss.

51. Milestone 1 — focal loss + gradcheck

Worked example

Your turn: implement focal_loss(logits, targets, gamma) and confirm it produces 0.2437 on the test batch. Then pass gradcheck.

Hint: p_t = p * targets + (1-p) * (1-targets); use F.binary_cross_entropy_with_logits(..., reduction='none').

import torch, torch.nn.functional as F
from torch.autograd import gradcheck

def focal_loss(logits, targets, gamma=2.0):
    p   = torch.sigmoid(logits)
    bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')
    p_t = p * targets + (1 - p) * (1 - targets)
    return ((1 - p_t) ** gamma * bce).mean()

logits  = torch.tensor([3.0, 0.1, -1.0, 0.5], requires_grad=True)
targets = torch.tensor([1.0, 0.0,  1.0, 1.0])
print(round(focal_loss(logits, targets).item(), 4))   # 0.2437

xf = torch.randn(5, dtype=torch.float64, requires_grad=True)
tf = torch.randint(0, 2, (5,)).double()
print(gradcheck(focal_loss, (xf, tf)))  # True
checkexpectedactual
mean focal loss0.24370.2437
gradcheckTrueTrue

52. What each one costs: Milestone 1 — focal loss + gradcheck

Trade off

Comparison matrix

From Milestone 1 — focal loss + gradcheck: every row here is a choice with a cost. Fill the actual column, then say which row you would actually pick and what you give up for it.

checkexpectedactual
mean focal loss0.24370.2437
gradcheckTrueTrue

53. Milestone 2 — triplet loss + hard mining

Worked example

Your turn: implement batch hard triplet loss. For each anchor, find the hardest positive and hardest negative in the batch using pairwise distances.

Hint: torch.cdist(emb, emb) gives all pairwise Euclidean distances. Use labels == labels[i] as a boolean mask.

import torch

def hard_triplet_loss(emb, labels, margin=1.0):
    n = len(emb)
    dists = torch.cdist(emb, emb)
    losses = []
    for i in range(n):
        pos_mask = (labels == labels[i]) & (torch.arange(n) != i)
        neg_mask = (labels != labels[i])
        if not pos_mask.any() or not neg_mask.any(): continue
        d_ap = dists[i][pos_mask].max()  # hardest positive
        d_an = dists[i][neg_mask].min()  # hardest negative
        losses.append(torch.clamp(d_ap - d_an + margin, min=0))
    return torch.stack(losses).mean()

torch.manual_seed(1)
emb = torch.randn(4, 8)
labels = torch.tensor([0, 0, 1, 1])
print(round(hard_triplet_loss(emb, labels).item(), 4))  # 1.8210
embedding paird(a,p) maxd(a,n) minloss
anchor=0 (class 0)3.83303.46901.3640
anchor=1 (class 0)3.83302.95001.8830
anchor=2 (class 1)4.22802.95002.2780
anchor=3 (class 1)4.22803.46901.7590

54. Which is which, by d(a,p) max

Discrimination

Sort into buckets

Sort these by d(a,p) max, from memory, without looking back at Milestone 2 — triplet loss + hard mining. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

3.8330
anchor=0 (class 0); anchor=1 (class 0)
4.2280
anchor=2 (class 1); anchor=3 (class 1)
g1
d(a,p) max is "3.8330" for anchor=0 (class 0), anchor=1 (class 0) — that is what the table on "Milestone 2 — triplet loss + hard mining" records, and it is the single property separating this group from the rest.
g2
d(a,p) max is "4.2280" for anchor=2 (class 1), anchor=3 (class 1) — that is what the table on "Milestone 2 — triplet loss + hard mining" records, and it is the single property separating this group from the rest.

55. The full program

Concept

import torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import make_classification

def focal_loss(logits, targets, gamma=2.0):
    p   = torch.sigmoid(logits)
    bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')
    p_t = p * targets + (1 - p) * (1 - targets)
    return ((1 - p_t) ** gamma * bce).mean()

torch.manual_seed(7)
X, y = make_classification(n_samples=500, n_features=8, n_informative=4,
                           weights=[0.9, 0.1], random_state=42)
Xt = torch.tensor(X, dtype=torch.float32)
yt = torch.tensor(y, dtype=torch.float32).unsqueeze(1)

for loss_fn, name in [(F.binary_cross_entropy_with_logits, 'CE'), (focal_loss, 'FL')]:
    mdl = nn.Sequential(nn.Linear(8,16), nn.ReLU(), nn.Linear(16,1))
    opt = torch.optim.Adam(mdl.parameters(), lr=1e-2)
    for _ in range(100):
        opt.zero_grad(); loss_fn(mdl(Xt), yt).backward(); opt.step()
    with torch.no_grad():
        pred = (torch.sigmoid(mdl(Xt)) > 0.5).float()
    recall = pred[yt.squeeze()==1].mean().item()
    print('%s minority recall: %.3f' % (name, recall))
lossminority recall
cross-entropy0.560
focal (gamma=2)0.680

If your focal minority recall exceeds CE's recall by at least 10 points — you've confirmed focal loss's core claim.

56. Fill in: minority recall for The full program

Comparison

Comparison matrix

From The full program: refill the minority recall column from what you know. The rest of the table is as it appeared.

lossminority recall
cross-entropy0.560
focal (gamma=2)0.680

57. Show it off

Concept

Out loud, slides closed: (1) Derive the focal loss modulating factor and explain why it helps class imbalance. (2) Describe what makes an 'online hard negative' and why easy triplets waste compute. (3) Explain the contract for a custom autograd.Function.

Stretch (homework): implement the alpha-balanced variant of focal loss (alpha * FL(p_t)); add L2 normalization to the triplet embeddings (unit sphere); implement contrastive loss and compare it to triplet loss on a 4-class toy dataset. Next: unsupervised deep learning — k-means, GMMs, and autoencoders (Lesson 52).

58. Connect it up: Lesson 51: Custom Loss Functions & Autograd

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — The differentiability requirement · Focal loss: down-weight the easy ones · Triplet loss: metric learning · Custom autograd.Function · Your turn: build the losses. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

59. What you can do now

Recap

losstaskkey idea
focalclass imbalance(1-p_t)^gamma down-weights easy examples
tripletmetric learning / facesmax(0, d_pos - d_neg + margin)
contrastivepair learning / CLIPyd^2 + (1-y)max(0, m-d)^2
custom Functionnon-standard opsctx.save_for_backward + chain rule

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 51 (Week 18 — Custom Loss Functions) — Barron · USAAIO Round 2 Preparation, 2026
  2. RetinaNet / Focal Loss paper: Lin et al., ICCV 2017 — Lin, T.-Y. et al. Focal Loss for Dense Object Detection. ICCV 2017.
  3. FaceNet: Schroff et al., CVPR 2015 (triplet loss) — Schroff, F. et al. FaceNet: A Unified Embedding for Face Recognition. CVPR 2015.
  4. All loss formulas, trace tables, and CE-vs-FL comparison verified with torch 2.7.1+cpu and sklearn, June 2026 — torch 2.7.1 + numpy 2.2.6 + scikit-learn, real execution, June 2026

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108