DFL:分布焦点损失——让框回归“学分布“,而不只是“猜数字“

出自论文 Generalized Focal Loss: Learning Qualified and Distributed Bounding Boxes for Dense Object Detection(CVPR 2020,Xiang Li et al.)。GFL 由 QFL(分类分支)与 DFL(回归分支)两部分组成。


一、解决什么问题

1. 传统做法的两个割裂

分支 传统做法 问题
分类 BCE + hard 0/1 标签 训练用 0/1,推理却用「分类分数 × IoU」打分,口径不一致
回归 L1 / IoU loss 监督一个标量 假设偏移服从狄拉克分布(只有一个确定值),表达不了遮挡、模糊、标注噪声下的定位不确定性

2. GFL 的解决思路

  • QFL :把分类标签从 hard 0/1 改成连续 IoU 软标签,让训练和推理的打分口径对齐。
  • DFL:把回归从「点估计」升级为「离散概率分布估计」,推理时对分布求期望(integral)还原连续偏移;分布的尖锐程度隐式表达定位质量。

二、QFL:Quality Focal Loss

1. 公式

设 σ\sigmaσ 为模型输出的分类分数(sigmoid 后),yyy 为预测框与 GT 的真实 IoU(软标签,∈0,1\in0,10,1):

QFL(σ)=−∣y−σ∣β(1−y)log⁡(1−σ)+ylog⁡σ,β=2QFL(\sigma) = -|y - \sigma|^{\beta}\big(1-y)\\log(1-\\sigma) + y\\log\\sigma\\big, \quad \beta = 2QFL(σ)=−∣y−σ∣β(1−y)log(1−σ)+ylogσ,β=2

  • (1−y)log⁡(1−σ)+ylog⁡σ(1-y)\log(1-\sigma) + y\log\sigma(1−y)log(1−σ)+ylogσ:以 IoU 为软标签的二分类交叉熵;
  • ∣y−σ∣β|y-\sigma|^{\beta}∣y−σ∣β:focal 调制项,预测越偏离真实 IoU,权重越大,聚焦难样本。

2. 关键点

  • 训练阶段不做「分类 × IoU」的显式相乘,相乘发生在推理打分阶段;
  • 正样本标签 = 该预测框与 GT 的 IoU(连续值),负样本标签 = 0;
  • β=2\beta=2β=2 是论文默认值。

三、DFL:Distribution Focal Loss

1. 离散化与分布建模

把连续偏移量 yyy 量化为 nbinsn_{\text{bins}}nbins 个离散候选。网络对每条边输出 nbinsn_{\text{bins}}nbins 个 logits,经 softmax 得到概率:

Sk=P(y=k),k=0,1,...,nbins−1S_k = P(y = k), \quad k = 0, 1, \dots, n_{\text{bins}}-1Sk=P(y=k),k=0,1,...,nbins−1

推理时对分布求期望(integral)还原连续偏移:

y^=∑k=0nbins−1k⋅Sk\hat{y} = \sum_{k=0}^{n_{\text{bins}}-1} k \cdot S_ky^=k=0∑nbins−1k⋅Sk

bin 索引与偏移量一一对应,期望即加权平均索引,天然支持小数------例如 7.4 由 bin 7 与 bin 8 的概率质量按比例插值表达。

2. DFL 公式

训练时真实偏移 yyy 是连续值,取其最近的两个整数 bin:

yi=⌊y⌋,yi+1=yi+1y_i = \lfloor y \rfloor, \qquad y_{i+1} = y_i + 1yi=⌊y⌋,yi+1=yi+1

DFL 只在这两个 bin 上做带权交叉熵:

DFL(Si,Si+1)=−(yi+1−y)log⁡Si+(y−yi)log⁡Si+1DFL(S_i, S_{i+1}) = -\big(y_{i+1} - y)\\log S_i + (y - y_i)\\log S_{i+1}\\bigDFL(Si,Si+1)=−(yi+1−y)logSi+(y−yi)logSi+1

记权重:

wl=yi+1−y,wr=y−yi,wl+wr=1w_l = y_{i+1} - y, \qquad w_r = y - y_i, \qquad w_l + w_r = 1wl=yi+1−y,wr=y−yi,wl+wr=1

则:

DFL=−wllog⁡Si+wrlog⁡Si+1DFL = -\bigw_l \\log S_i + w_r \\log S_{i+1}\\bigDFL=−wllogSi+wrlogSi+1

目标离哪个 bin 越近,该 bin 获得越大的监督权重。

3. 符号表

符号 含义
yyy 连续的真实偏移量(GT),例如 7.4
nbinsn_{\text{bins}}nbins 离散化粒度,候选 bin 数量
kkk bin 索引,k=0,...,nbins−1k=0,\dots,n_{\text{bins}}-1k=0,...,nbins−1
SkS_kSk 网络预测「偏移落在 bin kkk」的概率(softmax 后)
yi,yi+1y_i, y_{i+1}yi,yi+1 目标 yyy 最近的左、右两个整数 bin
wl,wrw_l, w_rwl,wr 监督权重,离哪个 bin 越近权重越大
y^\hat{y}y^ 推理输出的连续偏移量(分布期望 / integral)

4. 两个关键直觉

直觉一:让分布聚焦,加速收敛

仅监督「期望 + L1」时,同一个期望值可以对应无数种形态完全不同的概率分布,模型缺乏约束分布形态的梯度信号,收敛缓慢。DFL 直接把监督压在相邻两个 bin 上,迫使概率快速在正确位置聚成尖峰。

直觉二:只监督两个 bin,实现亚 bin 小数插值

若对全部 bin 做 one-hot 交叉熵(只监督 ⌊y⌋\lfloor y\rfloor⌊y⌋ 这一个 bin),目标的小数部分完全丢失。DFL 按到 bin 的距离分配监督权重,让模型学会用相邻两个 bin 的概率质量比例编码任意连续值。

5. 与 label smoothing 的区别

label smoothing DFL
概率去向 从真类让出 ε\varepsilonε,均摊到全部类别 只压到最近的两个 bin,其余 bin 无直接监督、靠 softmax 竞争被压低
目的 防过拟合、提升泛化 逼分布又尖又准、保留小数插值信息

四、PyTorch 工程实现

先截断连续 target,再计算 floor 和权重,避免越界时产生负权重。

python 复制代码
import torch
import torch.nn.functional as F


def dfl(logits: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
    """
    Distribution Focal Loss 工程实现
    logits: [N, n_bins] 原始输出,未过 softmax
    target: [N,] 连续偏移 GT,合法范围 [0, n_bins - 1]
    """
    n_bins = logits.shape[-1]

    # 关键:先截断连续浮点 target,再求 floor,避免越界产生负权重
    target = torch.clamp(target, 0.0, float(n_bins - 1))
    y_i = target.floor().long()
    # 兜底:target 恰好等于 n_bins-1 时,y_i+1 不能越界
    y_i = torch.clamp(y_i, 0, n_bins - 2)
    y_i1 = y_i + 1

    w_right = target - y_i.float()   # w_r = y - y_i
    w_left = 1.0 - w_right           # w_l = y_{i+1} - y

    ce_left = F.cross_entropy(logits, y_i, reduction="none")
    ce_right = F.cross_entropy(logits, y_i1, reduction="none")

    loss = w_left * ce_left + w_right * ce_right
    return loss.mean()


def qfl(pred: torch.Tensor, target_iou: torch.Tensor,
        beta: float = 2.0, eps: float = 1e-6) -> torch.Tensor:
    """
    Quality Focal Loss
    pred:       [N,] sigmoid 输出的分类分数,范围 [0, 1]
    target_iou: [N,] 软标签,预测框与 GT 的真实 IoU,范围 [0, 1]
    beta:       focal 调制指数,论文默认 2
    eps:        防止 sigmoid 饱和时 log(0) 数值爆炸
    """
    # 数值稳定保护,避免 pred 趋近 0 或 1 产生 inf
    pred = torch.clamp(pred, eps, 1.0 - eps)

    ce = F.binary_cross_entropy(pred, target_iou, reduction="none")
    scale = torch.abs(target_iou - pred).pow(beta)
    loss = scale * ce
    return loss.mean()


def integral_projection(logits: torch.Tensor) -> torch.Tensor:
    """
    推理阶段:由 logits 求期望,得到连续偏移值(integral)
    logits: [N, n_bins]
    return: [N,]
    """
    probs = F.softmax(logits, dim=-1)
    bin_idx = torch.arange(probs.shape[-1], device=probs.device, dtype=probs.dtype)
    return (probs * bin_idx).sum(dim=-1)


if __name__ == "__main__":
    # ---- DFL 测试(含边界 case)----
    N, n_bins = 3, 16
    logits = torch.randn(N, n_bins)
    dfl_target = torch.tensor([7.4, 2.2, 15.9])  # 15.9 是越界边界 case
    print("dfl loss:", dfl(logits, dfl_target).item())
    print("pred offset:", integral_projection(logits))

    # ---- QFL 测试 ----
    pred_score = torch.sigmoid(torch.randn(N))
    iou_gt = torch.tensor([0.85, 0.3, 0.0])
    print("qfl loss:", qfl(pred_score, iou_gt).item())

六、小结

GFL 双损失分工一句话:QFL 解决分类分支「训练期 hard 标签、推理期乘 IoU」的口径不一致;DFL 解决回归分支「点估计」无法表达分布与不确定度的缺陷,二者共同构成 GFL。

DFL 本质:定向软标签交叉熵,仅监督 GT 相邻两个 bin,按距离分配权重。

DFL 三大作用

  1. 加速收敛:概率快速聚焦,分布锐化;
  2. 亚 bin 精度:两 bin 概率比例编码小数;
  3. 隐式表达定位不确定度:分布越尖锐定位越可信。

局限

  • bin 数量超参敏感:真实偏移超出 bin 范围会被截断,精度下降;增大 bin 数量带来通道数膨胀;
  • 仅监督相邻两 bin,远处 bin 依靠 softmax 竞争与 IoU 损失间接约束;
  • 四边回归互相独立,本身不约束框几何关系,必须搭配 IoU 类损失(如 GIoU / CIoU);
  • 回归分支通道数上升,带来少量计算开销。

笔记用于学习理解,上线请参考官方实现。

相关推荐
唐兴通个人1 小时前
新华保险集团携手浙江大学,邀请唐兴通老师主讲AI时代b保险新媒体营销增长专项培训
人工智能
AIGC大时代1 小时前
文献驱动学发现:Scimon、Scideator与CHIMERA 怎么验收(Tom Hope)
人工智能·科技·机器学习
IT古董1 小时前
AI资讯日报|2026年9月7日:涉 AI 纠纷裁判规则出台,GPT-6 引领新一波能力升级,国产算力与智能体加速落地
人工智能
jimmyleeee2 小时前
大模型安全之六:LLM过度代理(Excessive Agency)
人工智能·安全
UCloud_TShare2 小时前
优刻得孔明智算平台携手openFuyao,加速多样化算力规模化落地
人工智能·ai·大模型
染指11102 小时前
110.Agent-LangChain核心组件-Messages消息和提示词工程
人工智能·microsoft·langchain·agents
一航jason2 小时前
GPU 推荐用吗?——8295 上 Adreno 695 的现实评估
人工智能·ai·ai编程·llama·ai-native
光锥智能2 小时前
小鹏第一位机器人自己走下产线
大数据·人工智能
YonyouHRSaaS2 小时前
AI视频面试系统定义、功能作用、品牌推荐、选择攻略
人工智能·面试·职场和发展·ai面试·视频面试·ai视频面试