4.2e — Détection d’objets from scratch : la Focal Loss

← 4.2c — Détection anchor-based from scratch · Série 04-Vision

Le notebook 4.2c a construit la chaîne complète d’un détecteur anchor-based dont le déséquilibre pos:neg est géré par sous-échantillonnage (ratio 1:3). Ce notebook aborde la troisième voie, proposée par Lin et al. (RetinaNet, 2017) : redéfinir la loss pour qu’elle traite elle-même le déséquilibre, sans sous-échantillonnage.

L’idée centrale : dans le régime extrême (≈ 1000 negatives pour 1 positive), la binary cross-entropy est dominée par les easy negatives — ceux que le classifieur classe déjà correctement avec une probabilité \(\approx 0\). Leur gradient reste faible mais leur nombre cumulé noie complètement le signal des positifs. La focal loss multiplie chaque contribution par un facteur \((1 - p_t)^\gamma\) qui annule les easy examples et préserve les hard examples — d’où le nom.

On construit ici :

Aucun torchvision.ops.sigmoid_focal_loss ni focal_loss d’une lib tierce — on l’écrit, et on l’instrumente, pour mesurer ce qu’elle fait réellement.

1. Le problème : easy negatives qui dominent

La binary cross-entropy standard est :

\[\text{BCE}(p, y) = -y \log p - (1 - y) \log(1 - p)\]

où \(y \in \{0, 1\}\) est la classe cible et \(p = \sigma(z)\) la probabilité prédite. Pour un easy negative (\(y = 0\), \(p \approx 0{,}01\)), la contribution à la loss est \(\approx -\log(0{,}99) \approx 0{,}01\) — très petite individuellement. Mais sur 1000 easy negatives par image, leur contribution cumulée est \(\approx 10\), soit \(10^4\) fois la contribution d’un positif seul (\(\approx 0{,}001\)). Le gradient total est dominé par les easy negatives ; le réseau n’apprend rien sur les positifs parce qu’ils sont statistiquement invisibles.

Le remède classique est numérique : sous-échantillonner les négatifs pour rééquilibrer (notebook 4.2c, ratio 1:3). Le remède de RetinaNet est analytique : redéfinir la loss pour qu’elle-même pondère moins les easy examples. C’est l’objet de la section suivante.

import math
import time

import matplotlib.pyplot as plt
import numpy as np
import torch
import os
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")  # determinisme cuBLAS (#16795)
# Determinisme (#16795) : la graine seule ne garantit PAS la reproductibilite
# (heuristiques cuDNN, kernels non deterministes). warn_only=True au premier
# passage pour inventorier les ops fautives sans faire echouer le run.
torch.use_deterministic_algorithms(True, warn_only=True)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
import torch.nn as nn
import torch.nn.functional as F

SEED = 0
torch.manual_seed(SEED)
np.random.seed(SEED)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device:", DEVICE.type)
device: cuda

2. Dérivation de la focal loss depuis la BCE pondérée

Définissons \(p_t\) comme la probabilité de la classe cible :

\[p_t = \begin{cases} p & \text{si } y = 1 \\ 1 - p & \text{si } y = 0 \end{cases}\]

Avec cette convention, BCE se récrit \(\text{BCE}(p_t) = -\log p_t\). La BCE pondérée par \(\alpha \in [0, 1]\) est \(\text{BCE}_\alpha(p_t) = -\alpha_t \log p_t\) où \(\alpha_t = \alpha\) si \(y = 1\) et \(\alpha_t = 1 - \alpha\) sinon. C’est le cas particulier \(\gamma = 0\) de la famille :

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

Le facteur \((1 - p_t)^\gamma\) est le modulateur :

  • si \(p_t\) est proche de 1 (example bien classé), \((1 - p_t)^\gamma \approx 0\) ⇒ sa contribution à la loss s’effondre ;
  • si \(p_t\) est proche de 0 (example mal classé), \((1 - p_t)^\gamma \approx 1\) ⇒ sa contribution reste pleine.

Le paramètre \(\gamma\) (focusing parameter, \(\gamma \geq 0\)) règle l’agressivité de l’écrasement. \(\gamma = 0\) redonne la BCE pondérée ; \(\gamma = 2\) est le choix par défaut de RetinaNet. \(\alpha_t\) équilibre le ratio pos:neg indépendamment du \(\gamma\).

def focal_loss(logits, targets, gamma=2.0, alpha=0.25, reduction="mean"):
    """Focal loss from scratch, vectorisée.

    logits   : (N,) ou (N, C) — logits bruts (avant sigmoid)
    targets  : (N,) ou (N, C) — 0 ou 1 (float)
    gamma    : focusing parameter (>=0, défaut 2.0 comme RetinaNet)
    alpha    : poids de la classe positive (défaut 0.25 pour pos:neg ≈ 1:3)
    reduction : "mean" | "sum" | "none"

    Implementation :
    p   = sigmoid(logits)                # probabilité prédite pour la classe 1
    p_t = p * targets + (1 - p) * (1 - targets)  # proba de la classe cible
    alpha_t = alpha * targets + (1 - alpha) * (1 - targets)
    loss = - alpha_t * (1 - p_t) ** gamma * log(p_t)

    Compatible multi-classe : on aplatit logits et targets en (N*C,).
    """
    flat_logits = logits.reshape(-1)
    flat_targets = targets.reshape(-1).to(flat_logits.dtype)
    p = torch.sigmoid(flat_logits)
    p_t = p * flat_targets + (1.0 - p) * (1.0 - flat_targets)
    alpha_t = alpha * flat_targets + (1.0 - alpha) * (1.0 - flat_targets)
    eps = 1e-9
    loss = -alpha_t * (1.0 - p_t).pow(gamma) * torch.log(p_t.clamp(min=eps))
    if reduction == "mean":
        return loss.mean()
    if reduction == "sum":
        return loss.sum()
    return loss


# Vérifications : cas limites
# 1) gamma=0 -> BCE ponderee par alpha. Avec alpha=0.5, chaque exemple
# recoit le meme poids (alpha_t = 0.5), donc FL(gamma=0, alpha=0.5) = 0.5 * BCE(mean).
logits = torch.tensor([2.0, -2.0, 0.0])
targets = torch.tensor([1.0, 0.0, 1.0])
fl_g0 = focal_loss(logits, targets, gamma=0.0, alpha=0.5)
bce_mean = F.binary_cross_entropy_with_logits(logits, targets, reduction="mean")
bce_weighted = 0.5 * bce_mean   # alpha=0.5 -> chaque exemple recoit poids 0.5
print(f"FL(gamma=0, alpha=0.5) = {fl_g0.item():.6f}")
print(f"0.5 * BCE(mean)        = {bce_weighted:.6f}")
print(f"écart attendu ≈ 0 :     {abs(fl_g0.item() - bce_weighted) < 1e-5}")

# 2) sur un example bien classé (p_t proche de 1) la FL tend vers 0
logits = torch.tensor([5.0])         # sigmoid -> 0.993
targets = torch.tensor([1.0])        # bien classé
print(f"FL(p_t=0.993, y=1) = {focal_loss(logits, targets).item():.6e}  (quasi nul)")

# 3) sur un example mal classé (p_t proche de 0) la FL reste pleine
logits = torch.tensor([-5.0])        # sigmoid -> 0.007
targets = torch.tensor([1.0])        # mal classé
print(f"FL(p_t=0.007, y=1) = {focal_loss(logits, targets).item():.6e}  (plein régime)")
FL(gamma=0, alpha=0.5) = 0.157834
0.5 * BCE(mean)        = 0.157834
écart attendu ≈ 0 :     True
FL(p_t=0.993, y=1) = 7.520144e-08  (quasi nul)
FL(p_t=0.007, y=1) = 1.234980e+00  (plein régime)

3. Visualisation : CE vs focal loss en fonction de \(p_t\)

Avant de comparer les gradients, regardons la loss elle-même. Pour \(\alpha = 0{,}25\) et \(\gamma \in \{0, 1, 2, 5\}\), on trace la contribution d’un example à la loss en fonction de \(p_t\). Sur la classe positive (\(y = 1\), donc \(p_t = p\)) :

  • BCE (\(\gamma = 0\)) : la loss reste non-nulle même pour \(p_t\) proche de 1 (l’easy positive côté \(y = 1\) se traduit par \(p_t\) proche de 1 ; symétriquement, l’easy negative \(y = 0\) a \(p_t = 1 - p \approx 1\), et la BCE reste positive partout).
  • Focal avec \(\gamma \geq 1\) : la loss s’effondre quand \(p_t \to 1\), et reste significative quand \(p_t\) est petit. C’est précisément ce qui rééquilibre le gradient total dans le régime déséquilibré.

À \(p_t = 0{,}9\) : la CE est \(\approx -0{,}25 \log(0{,}9) \approx 0{,}026\) ; la focal (\(\gamma = 2\)) est \(\approx -0{,}25 \cdot 0{,}01 \cdot \log(0{,}9) \approx 2{,}6 \times 10^{-4}\), soit 100× plus petite. À \(p_t = 0{,}1\) (mal classé) : la focal ne réduit la loss que d’un facteur \(\approx (1 - 0{,}1)^2 \approx 0{,}81\) — l’écart est marginal, c’est exactement le but.

p_t = np.linspace(0.01, 0.99, 200)
alpha = 0.25
fig, axes = plt.subplots(1, 2, figsize=(12, 4))
for gamma in [0, 1, 2, 5]:
    fl = -alpha * (1 - p_t) ** gamma * np.log(p_t)
    axes[0].plot(p_t, fl, label=f"γ={gamma}")
axes[0].set_xlabel("$p_t$ (probabilité de la classe cible)")
axes[0].set_ylabel("FL")
axes[0].set_title("Focal loss vs $p_t$ (α=0.25, classe positive)")
axes[0].legend(); axes[0].grid(alpha=0.3); axes[0].set_yscale("log")

# ratio FL / BCE : combien de fois la focal est plus petite que la CE
for gamma in [1, 2, 5]:
    ratio = (1 - p_t) ** gamma
    axes[1].plot(p_t, ratio, label=f"γ={gamma}")
axes[1].axhline(1.0, color="black", linestyle=":", label="BCE")
axes[1].set_xlabel("$p_t$")
axes[1].set_ylabel("(1 - $p_t$)$^γ$ = FL / BCE")
axes[1].set_title("Combien de fois la focal est plus petite que la BCE")
axes[1].legend(); axes[1].grid(alpha=0.3); axes[1].set_yscale("log")
plt.tight_layout(); plt.show()

4. La preuve du déséquilibre : CE vs focal sur 1000:1

Construisons un mini-batch pathologique : 1 positive et 1000 negatives, dont 950 sont des easy negatives (\(p \approx 0{,}02\) car le classifieur les classe déjà bien comme négatives) et 50 sont des hard negatives (\(p \approx 0{,}6\), le classifieur hésite). On regarde la somme des losses sur ce batch :

  • en CE : les 50 hard negatives dominent déjà (\(\approx 50 \times -\log 0{,}4 = 45{,}6\)) parce qu’elles sont moins bien classées que les easy neg (\(- \log 0{,}98 = 0{,}02\)) — la BCE est dominée par la difficulté individuelle, pas par le nombre d’exemples. Les easy neg contribuent \(\approx 950 \times 0{,}02 = 19{,}0\) (29 % du total) ; le positif \(\approx 0{,}1\). Total mesuré \(\approx 64{,}8\).
  • en focal (\(\gamma = 2\)) : les easy negatives voient leur contribution multipliée par \((1 - 0{,}98)^2 = 4 \times 10^{-4}\) — leur somme passe de 19,0 à \(0{,}004\) (0,0 % du total FL). Les hard negatives sont multipliées par \((1 - 0{,}6)^2 = 0{,}16\) ⇒ leur somme passe de 45,6 à 8,18. Le positif est multiplié par \((1 - 0{,}9)^2 = 0{,}01\) ⇒ 0,001.

Résultat mesuré cellule 9 : BCE total = 64,8 (29 % easy neg / 70 % hard neg / 0,2 % pos) contre FL(γ=2) total = 8,2 (0,0 % easy neg / 99,9 % hard neg / 0,0 % pos). La différence-clé n’est pas le volume (les deux cas sont dominés par les hard neg, ce qui est correct en classification) — c’est que la focal loss élimine entièrement la contribution des easy negatives, libérant le signal pour qu’il porte uniquement sur les exemples que le classifieur n’a pas encore maîtrisés. C’est l’effet recherché : un classifieur qui n’apprend plus des choses déjà apprises.

torch.manual_seed(7)
N_EASY = 950
N_HARD = 50
N_POS = 1

# logits générés : easy negatives très négatifs (p ≈ 0.02), hard negatives moyens (p ≈ 0.6), positif très positif
easy_neg_logits = torch.full((N_EASY,), -3.9)         # sigmoid(−3.9) ≈ 0.020
hard_neg_logits = torch.full((N_HARD,), 0.4)           # sigmoid(0.4) ≈ 0.598
pos_logits = torch.full((N_POS,), 2.2)                 # sigmoid(2.2) ≈ 0.901
logits = torch.cat([easy_neg_logits, hard_neg_logits, pos_logits])
targets = torch.cat([torch.zeros(N_EASY + N_HARD), torch.ones(N_POS)])

# BCE standard (alpha=1 sur les deux classes)
bce_per = F.binary_cross_entropy_with_logits(logits, targets, reduction="none")
# Focal loss gamma=2, alpha=0.5 (équilibré pour cette preuve)
fl_per = focal_loss(logits, targets, gamma=2.0, alpha=0.5, reduction="none")

bce_easy, bce_hard, bce_pos = bce_per[:N_EASY].sum(), bce_per[N_EASY:N_EASY + N_HARD].sum(), bce_per[N_EASY + N_HARD:].sum()
fl_easy, fl_hard, fl_pos = fl_per[:N_EASY].sum(), fl_per[N_EASY:N_EASY + N_HARD].sum(), fl_per[N_EASY + N_HARD:].sum()

bce_total = bce_easy + bce_hard + bce_pos
fl_total = fl_easy + fl_hard + fl_pos

print(f"BCE  total = {bce_total.item():7.3f}  | easy {bce_easy.item():7.3f} ({100*bce_easy/bce_total:5.1f}%)  hard {bce_hard.item():7.3f} ({100*bce_hard/bce_total:5.1f}%)  pos {bce_pos.item():7.3f} ({100*bce_pos/bce_total:5.1f}%)")
print(f"FL(γ=2) total = {fl_total.item():7.3f}  | easy {fl_easy.item():7.3f} ({100*fl_easy/fl_total:5.1f}%)  hard {fl_hard.item():7.3f} ({100*fl_hard/fl_total:5.1f}%)  pos {fl_pos.item():7.3f} ({100*fl_pos/fl_total:5.1f}%)")
BCE  total =  64.794  | easy  19.038 ( 29.4%)  hard  45.651 ( 70.5%)  pos   0.105 (  0.2%)
FL(γ=2) total =   8.186  | easy   0.004 (  0.0%)  hard   8.181 ( 99.9%)  pos   0.001 (  0.0%)

5. Comparaison d’entraînement : CE vs focal sur un problème déséquilibré

Construisons un classifieur binaire volontairement pathologique :

  • des features 2D \(x \in \mathbb{R}^2\) distribuées en deux gaussiennes (classe 0 large, classe 1 petit cluster isolé) ;
  • un MLP à deux couches, entraîné pendant 30 époques avec Adam (\(lr = 1\mathrm{e}{-2}\)) ;
  • deux entraînements : un avec BCE standard, un avec focal (\(\gamma = 2\), \(\alpha = 0{,}5\)) ;
  • on suit la convergence de la loss et de l’accuracy par époque, ainsi que le gradient sur le dernier batch (norme, distribution entre classes).

Hypothèse falsifiable : la focal loss doit converger plus rapidement vers l’accuracy 1.0 sur les positifs que la BCE, qui passe l’essentiel de ses premières époques à “détasser” les easy negatives.

def make_dataset(n_pos=80, n_neg=8000, seed=0):
    rng = np.random.default_rng(seed)
    X_pos = rng.normal(loc=[2.0, 2.0], scale=0.4, size=(n_pos, 2))
    X_neg = rng.normal(loc=[-1.0, -1.0], scale=1.5, size=(n_neg, 2))
    X = np.vstack([X_pos, X_neg]).astype(np.float32)
    y = np.hstack([np.ones(n_pos), np.zeros(n_neg)]).astype(np.float32)
    perm = rng.permutation(len(X))
    return torch.tensor(X[perm]), torch.tensor(y[perm])


Xtr, ytr = make_dataset(seed=1)
Xva, yva = make_dataset(seed=2)
print(f"train : {len(Xtr)} samples, {int(ytr.sum())} positifs ({100*ytr.mean():.2f}%)")
print(f"val   : {len(Xva)} samples, {int(yva.sum())} positifs ({100*yva.mean():.2f}%)")


class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(2, 32), nn.ReLU(),
            nn.Linear(32, 1)
        )

    def forward(self, x):
        return self.net(x).squeeze(-1)


def loss_share_per_class(model, X, y, loss_kind, gamma=2.0, alpha=0.5):
    """Part de la loss totale imputable aux positifs vs negatifs (batch).

    Pour chaque classe c : mean(loss_c) * |c| = perte non-réduite cumulée
    sur la classe. Total = somme sur les deux classes. Le partage par
    classe est (mean(loss_c) * |c|) / total — il reflète la
    **contribution absolue** de chaque classe à la loss totale,
    pondérée par la cardinalité. Avec 80 positifs et 8000 négatifs,
    un partage 50/50 signifierait mean(loss_pos) ≈ 100 * mean(loss_neg),
    ce qui est l'information recherchée.
    """
    model.eval()
    X = X.to(DEVICE); y = y.to(DEVICE)
    idx_pos = (y > 0.5).nonzero(as_tuple=True)[0]
    idx_neg = (y < 0.5).nonzero(as_tuple=True)[0]
    n_pos = len(idx_pos)
    n_neg = len(idx_neg)
    with torch.no_grad():
        if loss_kind == "bce":
            mean_pos = float(F.binary_cross_entropy_with_logits(model(X[idx_pos]), y[idx_pos]).item())
            mean_neg = float(F.binary_cross_entropy_with_logits(model(X[idx_neg]), y[idx_neg]).item())
        else:
            mean_pos = float(focal_loss(model(X[idx_pos]), y[idx_pos], gamma=gamma, alpha=alpha).item())
            mean_neg = float(focal_loss(model(X[idx_neg]), y[idx_neg], gamma=gamma, alpha=alpha).item())
    sum_pos = mean_pos * n_pos
    sum_neg = mean_neg * n_neg
    total = sum_pos + sum_neg
    return sum_pos, sum_neg, (sum_pos / total if total > 0 else 0.0)


def train_one(loss_kind, epochs=30, lr=1e-2, gamma=2.0, alpha=0.5):
    torch.manual_seed(42)
    model = MLP().to(DEVICE)
    opt = torch.optim.Adam(model.parameters(), lr=lr)
    hist_loss, hist_acc, hist_grad_norm = [], [], []
    for ep in range(epochs):
        model.train()
        opt.zero_grad()
        logits = model(Xtr.to(DEVICE))
        if loss_kind == "bce":
            loss = F.binary_cross_entropy_with_logits(logits, ytr.to(DEVICE))
        else:
            loss = focal_loss(logits, ytr.to(DEVICE), gamma=gamma, alpha=alpha)
        loss.backward()
        grad_sq = 0.0
        for p_ in model.parameters():
            if p_.grad is not None:
                grad_sq += float(p_.grad.detach().pow(2).sum())
        hist_grad_norm.append(grad_sq ** 0.5)
        opt.step()
        hist_loss.append(float(loss.detach()))
        model.eval()
        with torch.no_grad():
            preds = (torch.sigmoid(model(Xva.to(DEVICE))) > 0.5).float()
            tp = (preds * yva.to(DEVICE)).sum().item()
            npos = int(yva.sum().item())
            recall_pos = tp / max(npos, 1)
            acc = (preds == yva.to(DEVICE)).float().mean().item()
        hist_acc.append((acc, recall_pos))
    return hist_loss, hist_acc, hist_grad_norm, model


t0 = time.time()
bce_loss, bce_acc, bce_grad_norm, model_bce = train_one("bce", epochs=30)
fl_loss, fl_acc, fl_grad_norm, model_fl = train_one("focal", epochs=30, gamma=2.0, alpha=0.5)
print(f"2 entraînements × 30 époques en {time.time() - t0:.1f} s")

bce_lp, bce_ln, bce_lshare = loss_share_per_class(model_bce, Xtr, ytr, "bce")
fl_lp, fl_ln, fl_lshare = loss_share_per_class(model_fl, Xtr, ytr, "focal")
print(f"\nLoss share (modèles entraînés) :  BCE pos = {100*bce_lshare:.1f}%   Focal pos = {100*fl_lshare:.1f}%")

# Métrique gradient par classe — version honnête (cardinalité-pondéré + cosine).
# L'ancien `g_pos² / (g_pos² + g_neg²)` n'était PAS une part additive : la
# décomposition `||w_pos g_pos + w_neg g_neg||²` contient un cross term
# `2 w_pos w_neg <g_pos, g_neg>` qui peut être ± et faire basculer la "share".
# On reporte donc les deux énergies cardinalité-pondérées + la cosine similarité.
def _flat_grad(model):
    grads = [p_.grad.detach().reshape(-1) for p_ in model.parameters() if p_.grad is not None]
    return torch.cat(grads) if grads else torch.zeros(1)


def grad_class_metric(model, X, y, objective):
    """||w_pos g_pos||², ||w_neg g_neg||², cosine(g_pos, g_neg).

    Mesure cardinalité-pondérée : w_c = |c| / n_total. Le cross term
    `2 w_pos w_neg <g_pos, g_neg>` est explicitement hors du print car
    il peut basculer le signe de la "share".
    """
    X = X.to(DEVICE); y = y.to(DEVICE)
    idx_pos = (y > 0.5).nonzero(as_tuple=True)[0]
    idx_neg = (y < 0.5).nonzero(as_tuple=True)[0]
    n_pos, n_neg = len(idx_pos), len(idx_neg)
    n_total = n_pos + n_neg
    w_pos = n_pos / n_total
    w_neg = n_neg / n_total
    for p_ in model.parameters():
        if p_.grad is not None:
            p_.grad = None
    if n_pos > 0:
        objective(model(X[idx_pos]), y[idx_pos]).backward()
        g_pos = _flat_grad(model)
    else:
        g_pos = torch.zeros(0)
    for p_ in model.parameters():
        if p_.grad is not None:
            p_.grad = None
    if n_neg > 0:
        objective(model(X[idx_neg]), y[idx_neg]).backward()
        g_neg = _flat_grad(model)
    else:
        g_neg = torch.zeros(0)
    g_pos_norm = float(g_pos.norm().item()) if g_pos.numel() > 0 else 0.0
    g_neg_norm = float(g_neg.norm().item()) if g_neg.numel() > 0 else 0.0
    e_pos = (w_pos * g_pos_norm) ** 2
    e_neg = (w_neg * g_neg_norm) ** 2
    if g_pos.numel() > 0 and g_neg.numel() > 0:
        cos = float(torch.nn.functional.cosine_similarity(
            g_pos.unsqueeze(0), g_neg.unsqueeze(0)
        ).item())
    else:
        cos = 0.0
    return e_pos, e_neg, cos


bce_e_pos, bce_e_neg, bce_cos = grad_class_metric(
    model_bce, Xtr, ytr, lambda l, t: F.binary_cross_entropy_with_logits(l, t))
fl_e_pos, fl_e_neg, fl_cos = grad_class_metric(
    model_fl, Xtr, ytr, lambda l, t: focal_loss(l, t, gamma=2.0, alpha=0.5))
print(f"\nGradient par classe (cardinalité-pondéré, modèles entraînés) :")
print(f"  BCE : ||w_pos g_pos||² = {bce_e_pos:.6f}   ||w_neg g_neg||² = {bce_e_neg:.6f}   cos(g_pos, g_neg) = {bce_cos:+.3f}")
print(f"  FL  : ||w_pos g_pos||² = {fl_e_pos:.6f}   ||w_neg g_neg||² = {fl_e_neg:.6f}   cos(g_pos, g_neg) = {fl_cos:+.3f}")
print(f"  Note : cross term `2 w_pos w_neg <g_pos, g_neg>` est ± — `pos share` additive n existe pas.")
print(f"Grad norm epoch 0  :  BCE = {bce_grad_norm[0]:.3f}   Focal = {fl_grad_norm[0]:.3f}")
print(f"Grad norm epoch 29 :  BCE = {bce_grad_norm[-1]:.3f}   Focal = {fl_grad_norm[-1]:.3f}")

fig, axes = plt.subplots(1, 3, figsize=(15, 3.5))
axes[0].plot(bce_loss, label="BCE", color="indianred")
axes[0].plot(fl_loss, label="Focal (γ=2, α=0.5)", color="seagreen")
axes[0].set_xlabel("époque"); axes[0].set_ylabel("loss train")
axes[0].set_title("Convergence de la loss"); axes[0].legend(); axes[0].grid(alpha=0.3)

axes[1].plot([a[1] for a in bce_acc], label="BCE — rappel pos", color="indianred")
axes[1].plot([a[1] for a in fl_acc], label="Focal — rappel pos", color="seagreen")
axes[1].set_xlabel("époque"); axes[1].set_ylabel("rappel sur les positifs (val)")
axes[1].set_title("Détection des positifs — l'écart attendu")
axes[1].legend(); axes[1].grid(alpha=0.3)

axes[2].plot(bce_grad_norm, label="BCE", color="indianred")
axes[2].plot(fl_grad_norm, label="Focal", color="seagreen")
axes[2].set_yscale("log")
axes[2].set_xlabel("époque"); axes[2].set_ylabel("||grad|| (log)")
axes[2].set_title("Norme du gradient par époque (log)")
axes[2].legend(); axes[2].grid(alpha=0.3, which="both")
plt.tight_layout(); plt.show()

print(f"\nRappel final sur les positifs :  BCE = {bce_acc[-1][1]:.2f}   |   Focal = {fl_acc[-1][1]:.2f}")
print(f"Accuracy globale finale       :  BCE = {bce_acc[-1][0]:.3f}   |   Focal = {fl_acc[-1][0]:.3f}")
train : 8080 samples, 80 positifs (0.99%)
val   : 8080 samples, 80 positifs (0.99%)
2 entraînements × 30 époques en 1.8 s

Loss share (modèles entraînés) :  BCE pos = 46.9%   Focal pos = 32.5%

Gradient par classe (cardinalité-pondéré, modèles entraînés) :
  BCE : ||w_pos g_pos||² = 0.001455   ||w_neg g_neg||² = 0.002741   cos(g_pos, g_neg) = -0.596
  FL  : ||w_pos g_pos||² = 0.000031   ||w_neg g_neg||² = 0.000083   cos(g_pos, g_neg) = -0.943
  Note : cross term `2 w_pos w_neg <g_pos, g_neg>` est ± — `pos share` additive n existe pas.
Grad norm epoch 0  :  BCE = 1.189   Focal = 0.275
Grad norm epoch 29 :  BCE = 0.045   Focal = 0.005


Rappel final sur les positifs :  BCE = 0.00   |   Focal = 0.84
Accuracy globale finale       :  BCE = 0.990   |   Focal = 0.997

6. Dérivation analytique du gradient

Pour comprendre pourquoi la focal loss modifie le paysage du gradient, dérivons \(\partial \text{FL} / \partial z\) où \(z\) est le logit (avant sigmoid). Posons \(p = \sigma(z)\), \(p_t\) comme défini §2, et \(\alpha_t\) idem. La dérivation passe par :

  1. \(\partial \text{FL} / \partial p_t = \alpha_t (1 - p_t)^{\gamma - 1} \left[ \gamma p_t \log p_t + (p_t - 1) \right] / p_t\) (dérivation directe) ;
  2. \(\partial p_t / \partial z = \sigma'(z) \cdot (2y - 1) = p(1 - p) \cdot (2y - 1)\) (chain rule avec \(\partial p / \partial z = p(1 - p)\) et la définition \(p_t = y p + (1 - y)(1 - p)\)) ;
  3. composition : \(\partial \text{FL} / \partial z = \partial \text{FL} / \partial p_t \cdot p (1 - p) \cdot (2y - 1)\).

Le facteur \((1 - p_t)^{\gamma - 1}\) dans la dérivée explique pourquoi un easy example (\(p_t \to 1\), donc \((1 - p_t) \to 0\)) voit son gradient s’évanouir — c’est la dynamique de la focal loss, complémentaire à la dynamique de la loss elle-même. C’est l’objet de l’exercice 1.

Exercices

Exercice 1 — Gradient analytique vs autograd

Implémentez focal_loss_grad(z, y, gamma, alpha) qui retourne le gradient \(\partial \text{FL} / \partial z\) analytiquement, puis vérifiez sur 100 logits aléatoires qu’il coincide avec torch.autograd.grad à 1e-5 près. Le test est falsifiable : si les deux divergent, soit la dérivation est fausse, soit le graphe ne voit pas la même loss.

Indice : partez de focal_loss et appelez torch.autograd.grad(fl, z, create_graph=False)[0] ; utilisez torch.allclose(grad_auto, grad_analytique, atol=1e-5).

def focal_loss_grad(z, y, gamma=2.0, alpha=0.5):
    """Exercice 1 : gradient analytique de la focal loss par rapport au logit z.

    Retourne ∂FL/∂z de même shape que z.
    """
    # TODO etudiant
    pass


def test_grad(gamma=2.0, alpha=0.5, n=100):
    torch.manual_seed(0)
    z = torch.randn(n, requires_grad=True)
    y = (torch.rand(n) > 0.7).float()
    # gradient analytique
    g_ana = focal_loss_grad(z, y, gamma=gamma, alpha=alpha)
    # gradient autograd
    fl = focal_loss(z, y, gamma=gamma, alpha=alpha, reduction="sum")
    g_auto = torch.autograd.grad(fl, z, create_graph=False)[0]
    if g_ana is None:
        print("À implémenter")
        return
    ok = torch.allclose(g_auto, g_ana, atol=1e-5)
    print(f"γ={gamma} α={alpha} : autograd vs analytique | max |Δ| = {(g_auto - g_ana).abs().max():.2e}  |  match : {ok}")

test_grad()
print("Exercice 1 à compléter")
À implémenter
Exercice 1 à compléter

Exercice 2 — Focal loss multi-classe (par classes)

La formulation multi-classe de la focal loss remplace \(\alpha\) par un vecteur \(\alpha_c\) (un poids par classe) et garde le modulateur \((1 - p_t)^\gamma\) sur la probabilité de la classe cible. Implémentez focal_loss_multiclass(logits, targets, gamma, alpha_vec) où logits est de shape (N, C) et targets est de shape (N,) (indices de classe) ou (N, C) (one-hot).

Vérifiez sur un dataset multi-classes déséquilibré (5 classes, ratio 100:10:5:3:1) que la focal loss multi-classe donne un meilleur rappel sur les classes rares que la cross-entropy standard.

Indice : la formule est la même, avec $p_t = $ softmax(logits)[classe_cible]. Pour le gradient, partez de la cross-entropy classique et greffez le modulateur \((1 - p_t)^\gamma\).

def focal_loss_multiclass(logits, targets, gamma=2.0, alpha_vec=None):
    """Exercice 2 : focal loss multi-classe.

    logits   : (N, C)
    targets  : (N,) — indices entiers des classes cibles
    alpha_vec : (C,) — poids par classe (défaut : alpha uniforme = 1/C)
    """
    # TODO etudiant
    pass


print("Exercice 2 à compléter")
Exercice 2 à compléter

Exercice 3 — Mini-détecteur avec focal loss

Reprenez le modèle AnchorNet du notebook 4.2c (backbone + tête objectness) et remplacez sa BCE par focal_loss(gamma=2, alpha=0.25). Ré-entraînez sur le même terrain synthétique (2000 images train, 400 val) avec un déséquilibre délibéré : générez 10 négatifs par image pour chaque positif, plutôt que le ratio 1:3 actuel.

Hypothèse falsifiable : à déséquilibre aggravé, la focal loss doit converger en accuracy de validation sans avoir besoin du sous-échantillonnage 1:3, et le mAP final doit être au moins aussi bon.

Indice : importez AnchorNet depuis 4.2c (ou recopiez sa définition), utilisez le générateur de terrain synthétique de 4.2c, et remplacez uniquement la image_loss. Le notebook 4.2c fixe les hyperparamètres et le harnais de mesure (mAP VOC07 + VOC10).

print("Exercice 3 à compléter — reprendre AnchorNet du 4.2c et remplacer la BCE par focal_loss")
Exercice 3 à compléter — reprendre AnchorNet du 4.2c et remplacer la BCE par focal_loss

Conclusion

  • La focal loss ne change pas l’objectif de la classification : c’est toujours un classifieur binaire (ou multi-classe) qui apprend à discriminer foreground / background. Elle change la pondération par example en fonction de la difficulté prédite.
  • Le modulateur \((1 - p_t)^\gamma\) a deux effets conjoints : la loss des easy examples s’effondre, et leur gradient aussi. Les deux se conjuguent pour libérer le signal des hard examples.
  • La mesure conjointe (cellule focal11, section 5) sur le batch d’entraînement (80 positifs / 8000 négatifs, cardinalité-pondéré) donne, sur les modèles entraînés (et non sur des modèles fraîchement initialisés) : part de la loss totale imputable aux positifs = mean(loss_pos) × 80 / (mean(loss_pos) × 80 + mean(loss_neg) × 8000). Énergie gradient cardinalité-pondérée par classe (modèle entraîné, backward avec la loss correcte) : ||w_pos g_pos||² et ||w_neg g_neg||² (où w_c = |c| / n_total) plus la cosine similarité entre g_pos et g_neg — la « part additive » du gradient batch n’existe pas car ||w_pos g_pos + w_neg g_neg||² contient le cross term 2 w_pos w_neg <g_pos, g_neg>. Ce que la focal loss change mesurablement à l’entraînement, c’est la norme globale du gradient par époque (modulateur (1 - p_t)^γ) et l’alignement entre gradients de classes — pas un renversement trivial de partage. C’est exactement ce que RetinaNet utilise pour détecter des objets sur ~100 000 anchors par image échantillonnée, dans un régime exemplifié à ~1:1000 dans Lin et al. 2017. La mesure numérique précise du ratio n’est pas assertée comme telle dans le papier fondateur (Lin et al. 2017 cite ce ratio à titre illustratif du régime opératoire, sans le porter dans une table de chiffres), mais le constat opérationnel — les easy negatives dominent le signal de gradient sans un mécanisme de pondération — est l’objet central de la loss, et c’est précisément ce que le modulateur (1 - p_t)^γ corrige.
  • \(\alpha\) équilibre les classes, \(\gamma\) équilibre la difficulté. Les deux se règlent indépendamment ; \(\alpha = 0{,}25\), \(\gamma = 2\) est le choix par défaut de RetinaNet pour des objets rares sur fond générique.
  • Comparée à la sous-échantillonnage (4.2c, ratio 1:3), la focal loss garde tous les exemples : aucune information n’est jetée. C’est son avantage statistique ; sa limite est que les hyperparamètres \(\gamma\) et \(\alpha\) sont un choix à régler par dataset.

Références : Lin et al., Focal Loss for Dense Object Detection, ICCV 2017 (RetinaNet, papier fondateur) · He et al., Mask R-CNN, ICCV 2017 (utilise une softmax cross-entropy pour la classification des RoI et une BCE par pixel pour les masques — la focal loss n’y apparaît pas ; le mécanisme de focal loss est propre à RetinaNet et aux détecteurs one-stage denses).

Retour au sommet