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,1∈0,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);
  • 回归分支通道数上升,带来少量计算开销。

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

相关推荐
m4Rk_1 小时前
【论文阅读】Agent 记忆机制(83):Inside Out——用可演化 PersonaTree 构建 Agent 的核心长期记忆
论文阅读·人工智能·学习·开源·github
运行时异常1 小时前
【WMS 仓储系统集成 AI Agent 实战】第 8 讲(终篇):生产部署与并发安全——Semaphore 放进 Flux.defer 的坑,压测抓了一晚上
人工智能·安全
杨杨杨大侠1 小时前
Jev、Kev、Laya:决策模型怎么选,什么时候需要微调?
人工智能·python·agent
代码方舟1 小时前
零信任架构实战:基于天远人企关联构建自动化供应链金融网关
运维·人工智能·架构·自动化
虹科网络安全1 小时前
“顶会”看安全(十八):超越越狱:揭示由能力边界模糊引发的 LLM 应用安全风险
人工智能·安全
Maiko Star2 小时前
* LangChain 提示词模板详解:ChatPromptTemplate 的使用与高级特性
java·人工智能·langchain
鬓戈2 小时前
Rust 语言与 AI 应用生态调研及学习路径
人工智能·学习·rust
天远API2 小时前
零信任架构实战:基于天远人企关联构建自动化图谱网关
网络·人工智能·架构·自动化
爱吃提升2 小时前
文生视频模型发展趋势(2026)
人工智能·音视频
AIGS0012 小时前
什么是本体语义平台?和知识图谱、传统数据中台的区别在哪
人工智能·知识图谱·数据中台·ai问数·本体语义平台·智能数据中台