Focal Loss Function
Focal loss multiplies cross-entropy by a factor that fades to zero on easy examples, so a thousand background boxes stop drowning each rare object.
Why Does This Exist?
Dense detectors evaluate around 100,000 anchors per image with single-digit foreground. Plain cross-entropy sums over all of them, so easy background dominates the gradient and the model converges to predicting background everywhere. Practitioners patched this with sampling heuristics until the 2017 RetinaNet paper replaced the patch with a loss: scale each example's contribution by , where is the predicted probability on the true class. Easy examples fade automatically, no sampling code needed.
Think of It Like This
A volume knob wired to confusion
Imagine a classroom where every student shouts answers, and the teacher's attention is the gradient. Cross-entropy hands every shouter an equal microphone, so the thousand confident students chanting the obvious answer drown out the ten confused ones.
Focal loss wires each microphone to a confusion meter. Confident correct students fade to a whisper; struggling students stay loud. The teacher hears exactly the class that needs teaching.
How It Actually Works
The formula and the gamma knob
, with the standard setting and an optional class weight. Verified numbers show the shaping: at , cross-entropy is 0.1054 while focal loss is 0.0011, roughly 100 times smaller. At , cross-entropy is 1.2040 and focal loss 0.5900, only about 2 times smaller. Gamma controls the steepness: recovers plain cross-entropy, higher gamma crushes easy examples harder.
Where it applies beyond boxes
Any extreme imbalance responds: foreground-background in segmentation, rare fraud cases in classification, hard negatives in retrieval. The signature that focal loss fits is a model stuck predicting the majority class with a falling loss, easy examples voting down every correction.
Code
import torch
def focal_loss(logits: torch.Tensor, targets: torch.Tensor, gamma: float = 2.0) -> torch.Tensor: ce = torch.nn.functional.binary_cross_entropy_with_logits(logits, targets, reduction="none") p_t = torch.exp(-ce) return (((1 - p_t) ** gamma) * ce).mean()
easy = torch.tensor([3.0]) # model sure and right: p_t near 0.95hard = torch.tensor([0.0]) # model unsure: p_t = 0.5t = torch.tensor([1.])print(round(focal_loss(easy, t).item(), 5)) # -> 0.00011print(round(focal_loss(hard, t).item(), 4)) # -> 0.1733Watch Out For
Gamma tuned for COCO, shipped on mild imbalance
Gamma 2 on a 3-to-1 imbalance over-suppresses merely medium examples and stalls learning. The symptom is a loss that flatlines high with under-confident scores everywhere. Sweep gamma from 0.5 upward on validation, and confirm plain weighted cross-entropy is not enough first.
Focal loss masking broken labels
Down-weighting easy examples concentrates learning on the hardest ones, which include mislabeled data. The symptom is training that chases noise and validation that degrades late. Clean labels before sharpening focus, or cap the damage with label smoothing.
The Quick Version
- Focal loss scales cross-entropy by , fading easy examples from the gradient.
- Verified effect at : 100x smaller loss on easy cases, 2x on hard ones.
- It replaces sampling heuristics for extreme imbalance like 100,000-to-1 anchors.
- Gamma must be tuned per imbalance level, and clean labels matter more, not less.