非极大值抑制( NMS)算法 Python 代码实现
·
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 浮点数)。
代码逻辑
-
解析边界框坐标:
x1a, y1a, x2a, y2a = bbox_a x1b, y1b, x2b, y2b = bbox_b将两个边界框的坐标解包为左上角
(x1, y1)和右下角(x2, y2)。 -
计算边界框面积:
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,确保边界框有效。
- 计算每个边界框的宽度(
-
计算交集面积:
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 - x1和y2 - y1,同样使用torch.clamp确保非负。 - 交集面积为宽度 × 高度。
- 交集区域的左上角坐标为
-
计算 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 后保留的边界框集合。
代码逻辑
-
按置信度排序:
_, 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)。
- 使用
-
初始化结果列表:
res_idx_ls = list()- 创建空列表
res_idx_ls用于存储保留的边界框索引。
- 创建空列表
-
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,直接退出循环(因为无需进一步抑制)。
- 每次循环选择置信度最高的边界框(
-
抑制重叠边界框:
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,仅保留未被抑制的边界框索引。
- 取出当前最高置信度的边界框
-
返回结果:
bboxes_kept = bboxes[res_idx_ls] return bboxes_kept- 根据保留的索引
res_idx_ls,返回对应的边界框集合。
- 根据保留的索引
算法思想
NMS 是一种贪心算法,核心思想是:
- 按置信度从高到低排序所有边界框。
- 选择置信度最高的边界框,保留它。
- 计算该边界框与其他边界框的 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[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 过程:
- 排序:
probs = [0.8, 0.9, 0.5],排序后索引为[1, 0, 2](即先处理bboxes[1],再bboxes[0],最后bboxes[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]。
- 选择
- 第二次循环:
- 选择
bboxes[2](置信度 0.5),加入结果。 - 剩余边界框为空,退出循环。
- 选择
- 结果:保留
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.]])
代码优点与潜在优化
优点
- 清晰的逻辑:代码结构清晰,
calc_bbox_iou和non_maximum_supression函数职责明确,易于理解和维护。 - 使用 PyTorch:利用 PyTorch 的张量操作,适合深度学习任务中的边界框处理。
- 鲁棒性:通过
torch.clamp和assert确保边界框的有效性,防止无效输入导致错误。
潜在优化
- 向量化操作:
- 当前
non_maximum_supression使用循环逐一计算 IoU,效率较低。对于大量边界框,可通过向量化操作(例如批量计算 IoU)提高性能。 - 示例优化:将所有边界框的 IoU 计算为矩阵操作,减少循环。
- 当前
- 边界情况处理:
- 若输入
bboxes或probs为空,代码未显式处理,可能导致错误。可添加输入验证。
- 若输入
- 返回置信度:
- 当前仅返回保留的边界框,未返回对应的置信度。在实际应用中,可能需要同时返回
probs[res_idx_ls]。
- 当前仅返回保留的边界框,未返回对应的置信度。在实际应用中,可能需要同时返回
- 并行化:
- 对于大规模边界框集合,可利用 GPU 并行计算 IoU,进一步提升性能。
总结
该代码实现了目标检测中常用的 IoU 计算和 NMS 算法,适用于筛选重叠的边界框。calc_bbox_iou 函数通过计算交集和并集面积,精确计算两个边界框的 IoU。non_maximum_supression 函数基于贪心策略,按置信度排序并逐一抑制重叠边界框。示例代码展示了如何处理一组边界框和置信度,输出符合预期。代码逻辑清晰,适合小规模任务,但在大规模场景下可通过向量化等优化提升效率。
更多推荐


所有评论(0)