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
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.
Objectives
torch.autograd.Function with correct forward and backward methodsWarm-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.
Section
Part 1 of 4
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.
torch.* ops, never raw Python math on tensorstorch.clamp, torch.norm, torch.sigmoid, F.binary_cross_entropy_with_logits — all graph-tracked.mean() / .sum() before .backward())Verify with torch.autograd.gradcheck: numerically differentiates your function and compares it to autograd. A passing gradcheck means your gradients are correct.
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.
Section
Part 2 of 4
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.
| class | count | role in BCE loss |
|---|---|---|
| background (easy neg) | ~99 900 | dominates gradient |
| object (positive) | ~100 | signal 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.
Comparison
Comparison matrix
From The imbalance problem: refill the count column from what you know. The rest of the table is as it appeared.
| class | count | role in BCE loss |
|---|---|---|
| background (easy neg) | ~99 900 | dominates gradient |
| object (positive) | ~100 | signal drowned out |
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=1 | gamma=2 |
|---|---|---|---|
| 0.1 (hard) | 2.3026 | 2.0723 | 1.8651 |
| 0.5 (medium) | 0.6931 | 0.3466 | 0.1733 |
| 0.9 (easy) | 0.1054 | 0.0105 | 0.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.
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=1 | gamma=2 |
|---|---|---|---|
| 0.1 (hard) | 2.3026 | 2.0723 | 1.8651 |
| 0.5 (medium) | 0.6931 | 0.3466 | 0.1733 |
| 0.9 (easy) | 0.1054 | 0.0105 | 0.0011 |
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.
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).
| logit | target | p | (1-p_t)^2 | BCE | FL |
|---|---|---|---|---|---|
| 3.0 | 1 | 0.9526 | 0.0022 | 0.0486 | 0.0001 |
| 0.1 | 0 | 0.5250 | 0.2756 | 0.7444 | 0.2052 |
| -1.0 | 1 | 0.2689 | 0.5344 | 1.3133 | 0.7019 |
| 0.5 | 1 | 0.6225 | 0.1425 | 0.4741 | 0.0676 |
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.
Concept
10% positive class (50 positives, 450 negatives, make_classification). A 2-layer MLP trained for 100 epochs with Adam (lr=0.01).
| loss | overall acc | minority recall |
|---|---|---|
| cross-entropy | 0.948 | 0.560 |
| focal (gamma=2) | 0.960 | 0.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.
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.
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.
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.
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.
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.
Section
Part 3 of 4
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.
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.
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.
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.8620Easy 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.
| scenario | d(a,p) | d(a,n) | loss |
|---|---|---|---|
| easy negative (far) | 0.5831 | 2.5000 | 0.0000 |
| hard negative (close) | 0.5831 | 0.7211 | 0.8620 |
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.
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.
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.
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.
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 type | distance d | loss (m=1.0) |
|---|---|---|
| same class, d=0.2 | 0.20 | 0.0400 (= d²) |
| same class, d=0.8 | 0.80 | 0.6400 (= d²) |
| diff class, d=0.2 | 0.20 | 0.6400 (= (1-0.2)²) |
| diff class, d=1.5 | 1.50 | 0.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.
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 type | distance d | loss (m=1.0) |
|---|---|---|
| same class, d=0.2 | 0.20 | 0.0400 (= d²) |
| same class, d=0.8 | 0.80 | 0.6400 (= d²) |
| diff class, d=0.2 | 0.20 | 0.6400 (= (1-0.2)²) |
| diff class, d=1.5 | 1.50 | 0.0000 (already separated) |
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.
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.
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.
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.
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.math.log since it's just a scalar operation.; Compute p_t = p for all samples (regardless of target label).Section
Part 4 of 4
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.
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.
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.
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.
| x | sigma(x) | forward f(x) | grad df/dx |
|---|---|---|---|
| -1.0 | 0.2689 | -0.2689 | 0.0723 |
| 0.0 | 0.5000 | 0.0000 | 0.5000 |
| 1.0 | 0.7311 | 0.7311 | 0.9277 |
| 2.0 | 0.8808 | 1.7616 | 1.0918 |
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?
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.
p_t = p*y + (1-p)*(1-y) → FL = ((1-p_t)**gamma * BCE).mean()clamp(d(a,p) - d(a,n) + margin, min=0); mine hard negatives onliney*d^2 + (1-y)*clamp(m-d, min=0)^2; operates on pairsforward returns result + calls ctx.save_for_backward; backward returns grad_output * df/dxtorch.autograd.gradcheck with float64 inputs before shippingWhy: 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.
Pattern
p_t = p*y + (1-p)*(1-y) → FL = ((1-p_t)**gamma * BCE).mean()clamp(d(a,p) - d(a,n) + margin, min=0); mine hard negatives onliney*d^2 + (1-y)*clamp(m-d, min=0)^2; operates on pairsforward returns result + calls ctx.save_for_backward; backward returns grad_output * df/dxtorch.autograd.gradcheck with float64 inputs before shippingSorting
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.
autograd.Functionautograd.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.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.
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.
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:
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.
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.
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?
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.
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.
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.
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?
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.
Section
Project
Concept
Implement all three losses from scratch, verify gradients, and compare CE vs focal on an imbalanced dataset.
| # | milestone | tool |
|---|---|---|
| 1 | focal loss + gradcheck | sigmoid, BCE, autograd |
| 2 | triplet loss + hard mining | cdist, clamp, label masks |
| 3 | CE vs focal on imbalanced data | make_classification, Adam |
Build rules: implement each loss as a standalone function; pass float64 tensors to gradcheck; never call .item() inside the loss.
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.
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| check | expected | actual |
|---|---|---|
| mean focal loss | 0.2437 | 0.2437 |
| gradcheck | True | True |
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.
| check | expected | actual |
|---|---|---|
| mean focal loss | 0.2437 | 0.2437 |
| gradcheck | True | True |
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 pair | d(a,p) max | d(a,n) min | loss |
|---|---|---|---|
| anchor=0 (class 0) | 3.8330 | 3.4690 | 1.3640 |
| anchor=1 (class 0) | 3.8330 | 2.9500 | 1.8830 |
| anchor=2 (class 1) | 4.2280 | 2.9500 | 2.2780 |
| anchor=3 (class 1) | 4.2280 | 3.4690 | 1.7590 |
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.
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))| loss | minority recall |
|---|---|
| cross-entropy | 0.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.
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.
| loss | minority recall |
|---|---|
| cross-entropy | 0.560 |
| focal (gamma=2) | 0.680 |
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).
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.
Recap
torch.autograd.Function with correct forward + backward| loss | task | key idea |
|---|---|---|
| focal | class imbalance | (1-p_t)^gamma down-weights easy examples |
| triplet | metric learning / faces | max(0, d_pos - d_neg + margin) |
| contrastive | pair learning / CLIP | yd^2 + (1-y)max(0, m-d)^2 |
| custom Function | non-standard ops | ctx.save_for_backward + chain rule |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.