import torch

def calc_bbox_iou(bbox_a, bbox_b):
    """
        bbox: torch.Tensor, size [4]
        (x1, y1, x2, y2) 左上角和右下角坐标
    """
    # 计算两个输入bbox的面积
    x1a, y1a, x2a, y2a = bbox_a
    x1b, y1b, x2b, y2b = bbox_b
    wa = torch.clamp(x2a - x1a, min=0)
    ha = torch.clamp(y2a - y1a, min=0)
    area_a = wa * ha
    assert area_a > 0
    wb = torch.clamp(x2b - x1b, min=0)
    hb = torch.clamp(y2b - y1b, min=0)
    area_b = wb * hb
    assert area_b > 0
    # 求解输入bbox相交的框,并计算交集面积
    x1 = max(x1a, x1b)
    y1 = max(y1a, y1b)
    x2 = min(x2a, x2b)
    y2 = min(y2a, y2b)
    w = torch.clamp(x2 - x1, min=0)
    h = torch.clamp(y2 - y1, min=0)
    # 计算 IoU
    intersect = w * h
    union = area_a + area_b - intersect
    iou = intersect / union
    return iou.item()


def non_maximum_supression(bboxes, probs, iou_thr=0.5):
    """
    args:
        bboxes <torch.Tensor> size is [b, 4] 输入bbox集合
        probs <torch.Tensor> size is [b] 各个bbox对应的置信度
    ret:
        bboxes_kept <torch.Tensor> 经过NMS后保留下来的bbox
    """
    # 根据置信度排序,得到概率从大到小的indices
    _, indices = torch.sort(probs, descending=True)
    res_idx_ls = list()
    while indices.numel() > 0:
        idx = indices[0]
        # 每轮抑制的第一个元素直接加入
        # 因为已经确定与已保留的bbox无重叠
        res_idx_ls.append(idx.item())
        if indices.numel() < 2:
            break
        cur_bbox = bboxes[idx]
        no_remove_ls = list()
        # 用新加入的保留bbox对后面的进行抑制
        # 由于前面所有的保留bbox都对后面的元素进行过抑制,因此
        # 只需要在加入当前的bbox抑制一遍即可
        for i, test_idx in enumerate(indices[1:]):
            if calc_bbox_iou(bboxes[test_idx], cur_bbox) < iou_thr:
                # 不重叠的bbox被保留下来
                no_remove_ls.append(i)
        # 更新保留下来的bbox indices列表
        indices = indices[1:][no_remove_ls]
    # 最终的NMS后的结果bbox
    bboxes_kept = bboxes[res_idx_ls]
    return bboxes_kept


bbox_a = torch.Tensor([1, 1, 3, 3])
bbox_b = torch.Tensor([2, 2, 4, 4])

iou = calc_bbox_iou(bbox_a, bbox_b)
print("bbox_a 和 bbox_b 的 IoU 为:", iou)


bboxes = torch.Tensor([[1,1,30,30],
                       [1,1,31,31],
                       [1,1,40,40]])
probs = torch.Tensor([0.8, 0.9, 0.5])
bboxes_kept = non_maximum_supression(bboxes, probs, iou_thr=0.9)
print("NMS 前 bbox 集合 \n", bboxes)
print("NMS 前各个 bbox 对应概率 \n", probs)
print("=" * 24)
print("NMS 后 bbox 集合 \n", bboxes_kept)

代码主要涉及两个功能:计算两个边界框(bounding box, 简称bbox)的交并比(IoU, Intersection over Union),以及基于非极大值抑制(NMS, Non-Maximum Suppression)算法筛选边界框。代码使用PyTorch库实现,下面将逐部分解析代码的结构、功能、算法逻辑和示例输出。


1. calc_bbox_iou 函数

功能

calc_bbox_iou 函数用于计算两个边界框的交并比(IoU)。IoU 是目标检测中常用的指标,用于衡量两个边界框的重叠程度。IoU 的值在 [0, 1] 范围内,值越大表示两个边界框的重叠程度越高。

输入参数
  • bbox_a:PyTorch 张量,表示第一个边界框的坐标,形状为 [4],格式为 (x1, y1, x2, y2),其中 (x1, y1) 是左上角坐标,(x2, y2) 是右下角坐标。
  • bbox_b:PyTorch 张量,表示第二个边界框的坐标,格式同上。
返回值
  • 返回 IoU 的标量值(通过 .item() 转换为 Python 浮点数)。
代码逻辑
  1. 解析边界框坐标

    x1a, y1a, x2a, y2a = bbox_a
    x1b, y1b, x2b, y2b = bbox_b
    

    将两个边界框的坐标解包为左上角 (x1, y1) 和右下角 (x2, y2)

  2. 计算边界框面积

    wa = torch.clamp(x2a - x1a, min=0)
    ha = torch.clamp(y2a - y1a, min=0)
    area_a = wa * ha
    assert area_a > 0
    wb = torch.clamp(x2b - x1b, min=0)
    hb = torch.clamp(y2b - y1b, min=0)
    area_b = wb * hb
    assert area_b > 0
    
    • 计算每个边界框的宽度(wa, wb)和高度(ha, hb),使用 torch.clamp 确保宽度和高度非负(防止无效边界框)。
    • 面积计算公式为:宽度 × 高度。
    • 使用 assert 检查面积是否大于 0,确保边界框有效。
  3. 计算交集面积

    x1 = max(x1a, x1b)
    y1 = max(y1a, y1b)
    x2 = min(x2a, x2b)
    y2 = min(y2a, y2b)
    w = torch.clamp(x2 - x1, min=0)
    h = torch.clamp(y2 - y1, min=0)
    intersect = w * h
    
    • 交集区域的左上角坐标为 (max(x1a, x1b), max(y1a, y1b)),右下角坐标为 (min(x2a, x2b), min(y2a, y2b))
    • 交集区域的宽度和高度分别为 x2 - x1y2 - y1,同样使用 torch.clamp 确保非负。
    • 交集面积为宽度 × 高度。
  4. 计算 IoU

    union = area_a + area_b - intersect
    iou = intersect / union
    return iou.item()
    
    • 并集面积 = 边界框 A 面积 + 边界框 B 面积 - 交集面积。
    • IoU = 交集面积 / 并集面积。
    • 使用 .item() 返回标量值。
示例分析

在代码中,定义了两个边界框:

bbox_a = torch.Tensor([1, 1, 3, 3])
bbox_b = torch.Tensor([2, 2, 4, 4])
  • bbox_a:左上角 (1, 1),右下角 (3, 3),面积 = (3-1) × (3-1) = 4
  • bbox_b:左上角 (2, 2),右下角 (4, 4),面积 = (4-2) × (4-2) = 4
  • 交集区域:左上角 (max(1,2), max(1,2)) = (2,2),右下角 (min(3,4), min(3,4)) = (3,3),面积 = (3-2) × (3-2) = 1
  • 并集面积:4 + 4 - 1 = 7
  • IoU:1 / 7 ≈ 0.142857

运行结果:

iou = calc_bbox_iou(bbox_a, bbox_b)
print("bbox_a 和 bbox_b 的 IoU 为:", iou)

输出:

bbox_a 和 bbox_b 的 IoU 为: 0.1428571492433548

2. non_maximum_supression 函数

功能

non_maximum_supression 函数实现非极大值抑制(NMS)算法,用于在目标检测中去除重叠的边界框,保留置信度高且重叠程度低的边界框。

输入参数
  • bboxes:PyTorch 张量,形状为 [b, 4],表示一组边界框,每个边界框由 (x1, y1, x2, y2) 表示。
  • probs:PyTorch 张量,形状为 [b],表示每个边界框的置信度分数。
  • iou_thr:浮点数,表示 IoU 阈值,默认为 0.5。IoU 大于该阈值的边界框将被抑制。
返回值
  • bboxes_kept:PyTorch 张量,经过 NMS 后保留的边界框集合。
代码逻辑
  1. 按置信度排序

    _, indices = torch.sort(probs, descending=True)
    
    • 使用 torch.sort 按置信度从高到低排序,返回排序后的索引 indices
    • 例如,若 probs = [0.8, 0.9, 0.5],则 indices = [1, 0, 2](因为 0.9 > 0.8 > 0.5)。
  2. 初始化结果列表

    res_idx_ls = list()
    
    • 创建空列表 res_idx_ls 用于存储保留的边界框索引。
  3. NMS 主循环

    while indices.numel() > 0:
        idx = indices[0]
        res_idx_ls.append(idx.item())
        if indices.numel() < 2:
            break
    
    • 每次循环选择置信度最高的边界框(indices[0]),将其索引加入 res_idx_ls
    • 如果剩余边界框数量少于 2,直接退出循环(因为无需进一步抑制)。
  4. 抑制重叠边界框

    cur_bbox = bboxes[idx]
    no_remove_ls = list()
    for i, test_idx in enumerate(indices[1:]):
        if calc_bbox_iou(bboxes[test_idx], cur_bbox) < iou_thr:
            no_remove_ls.append(i)
    indices = indices[1:][no_remove_ls]
    
    • 取出当前最高置信度的边界框 cur_bbox
    • 对剩余边界框(indices[1:])逐一计算 IoU,若 IoU 小于阈值 iou_thr,则保留该边界框的索引(no_remove_ls)。
    • 更新 indices,仅保留未被抑制的边界框索引。
  5. 返回结果

    bboxes_kept = bboxes[res_idx_ls]
    return bboxes_kept
    
    • 根据保留的索引 res_idx_ls,返回对应的边界框集合。
算法思想

NMS 是一种贪心算法,核心思想是:

  1. 按置信度从高到低排序所有边界框。
  2. 选择置信度最高的边界框,保留它。
  3. 计算该边界框与其他边界框的 IoU,若 IoU 超过阈值,则抑制(移除)这些边界框。
  4. 重复上述步骤,直到所有边界框都被处理。
示例分析

在代码中,定义了边界框和置信度:

bboxes = torch.Tensor([[1,1,30,30],
                       [1,1,31,31],
                       [1,1,40,40]])
probs = torch.Tensor([0.8, 0.9, 0.5])
  • 边界框:
    • bboxes[0]:左上角 (1,1),右下角 (30,30),面积 = 29 × 29 = 841
    • bboxes[1]:左上角 (1,1),右下角 (31,31),面积 = 30 × 30 = 900
    • bboxes[2]:左上角 (1,1),右下角 (40,40),面积 = 39 × 39 = 1521
  • 置信度:[0.8, 0.9, 0.5]
  • IoU 阈值:iou_thr = 0.9

NMS 过程

  1. 排序probs = [0.8, 0.9, 0.5],排序后索引为 [1, 0, 2](即先处理 bboxes[1],再 bboxes[0],最后 bboxes[2])。
  2. 第一次循环
    • 选择 bboxes[1](置信度 0.9),加入结果。
    • 计算 bboxes[1]bboxes[0] 的 IoU:
      • 交集区域:左上角 (1,1),右下角 (min(31,30), min(31,30)) = (30,30),面积 = 29 × 29 = 841
      • 并集面积:900 + 841 - 841 = 900
      • IoU = 841 / 900 ≈ 0.9344 > 0.9,因此抑制 bboxes[0]
    • 计算 bboxes[1]bboxes[2] 的 IoU:
      • 交集区域:左上角 (1,1),右下角 (min(31,40), min(31,40)) = (31,31),面积 = 30 × 30 = 900
      • 并集面积:900 + 1521 - 900 = 1521
      • IoU = 900 / 1521 ≈ 0.5917 < 0.9,保留 bboxes[2]
    • 更新 indices = [2]
  3. 第二次循环
    • 选择 bboxes[2](置信度 0.5),加入结果。
    • 剩余边界框为空,退出循环。
  4. 结果:保留 bboxes[1]bboxes[2],即 [[1,1,31,31], [1,1,40,40]]

运行结果:

bboxes_kept = non_maximum_supression(bboxes, probs, iou_thr=0.9)
print("NMS 前 bbox 集合 \n", bboxes)
print("NMS 前各个 bbox 对应概率 \n", probs)
print("=" * 24)
print("NMS 后 bbox 集合 \n", bboxes_kept)

输出:

NMS 前 bbox 集合 
 tensor([[ 1.,  1., 30., 30.],
         [ 1.,  1., 31., 31.],
         [ 1.,  1., 40., 40.]])
NMS 前各个 bbox 对应概率 
 tensor([0.8000, 0.9000, 0.5000])
========================
NMS 后 bbox 集合 
 tensor([[ 1.,  1., 31., 31.],
         [ 1.,  1., 40., 40.]])

代码优点与潜在优化

优点
  1. 清晰的逻辑:代码结构清晰,calc_bbox_iounon_maximum_supression 函数职责明确,易于理解和维护。
  2. 使用 PyTorch:利用 PyTorch 的张量操作,适合深度学习任务中的边界框处理。
  3. 鲁棒性:通过 torch.clampassert 确保边界框的有效性,防止无效输入导致错误。
潜在优化
  1. 向量化操作
    • 当前 non_maximum_supression 使用循环逐一计算 IoU,效率较低。对于大量边界框,可通过向量化操作(例如批量计算 IoU)提高性能。
    • 示例优化:将所有边界框的 IoU 计算为矩阵操作,减少循环。
  2. 边界情况处理
    • 若输入 bboxesprobs 为空,代码未显式处理,可能导致错误。可添加输入验证。
  3. 返回置信度
    • 当前仅返回保留的边界框,未返回对应的置信度。在实际应用中,可能需要同时返回 probs[res_idx_ls]
  4. 并行化
    • 对于大规模边界框集合,可利用 GPU 并行计算 IoU,进一步提升性能。

总结

该代码实现了目标检测中常用的 IoU 计算和 NMS 算法,适用于筛选重叠的边界框。calc_bbox_iou 函数通过计算交集和并集面积,精确计算两个边界框的 IoU。non_maximum_supression 函数基于贪心策略,按置信度排序并逐一抑制重叠边界框。示例代码展示了如何处理一组边界框和置信度,输出符合预期。代码逻辑清晰,适合小规模任务,但在大规模场景下可通过向量化等优化提升效率。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐